use serde::ser::{
Impossible, Serialize, SerializeMap, SerializeSeq, SerializeStruct, SerializeStructVariant,
SerializeTuple, SerializeTupleStruct, SerializeTupleVariant, Serializer,
};
use std::fmt::{self, Display, Write as _};
#[derive(Debug, Clone, PartialEq)]
pub struct NonFiniteFloat {
pub path: String,
pub value: f64,
}
impl Display for NonFiniteFloat {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.path.is_empty() {
write!(f, "value must be finite, got {}", self.value)
} else {
write!(f, "{} must be finite, got {}", self.path, self.value)
}
}
}
impl std::error::Error for NonFiniteFloat {}
pub fn ensure_serialized_floats_are_finite<T>(value: &T) -> Result<(), NonFiniteFloat>
where
T: Serialize + ?Sized,
{
let mut walker = FloatWalker { path: String::new() };
match value.serialize(&mut walker) {
Ok(()) => Ok(()),
Err(WalkError::NonFinite(found)) => Err(found),
Err(WalkError::Custom(_)) => Ok(()),
}
}
#[derive(Debug)]
enum WalkError {
NonFinite(NonFiniteFloat),
Custom(String),
}
impl Display for WalkError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
WalkError::NonFinite(found) => Display::fmt(found, f),
WalkError::Custom(message) => f.write_str(message),
}
}
}
impl std::error::Error for WalkError {}
impl serde::ser::Error for WalkError {
fn custom<T: Display>(msg: T) -> Self {
WalkError::Custom(msg.to_string())
}
}
struct FloatWalker {
path: String,
}
impl FloatWalker {
fn push_field(&mut self, name: &str) -> usize {
let restore = self.path.len();
if !self.path.is_empty() {
self.path.push('.');
}
self.path.push_str(name);
restore
}
fn push_index(&mut self, index: usize) -> usize {
let restore = self.path.len();
write!(self.path, "[{index}]").expect("`String`'s `fmt::Write` impl is infallible");
restore
}
fn pop_to(&mut self, restore: usize) {
self.path.truncate(restore);
}
fn check(&self, value: f64) -> Result<(), WalkError> {
if value.is_finite() {
Ok(())
} else {
Err(WalkError::NonFinite(NonFiniteFloat {
path: self.path.clone(),
value,
}))
}
}
}
struct KeyRenderer;
impl KeyRenderer {
fn not_a_scalar(shape: impl Display) -> WalkError {
WalkError::Custom(format!("map key is not a scalar: {shape}"))
}
fn variant_key(name: &'static str, variant_index: u32, variant: &'static str) -> String {
if variant.is_empty() {
format!("{name}#{variant_index}")
} else {
variant.to_string()
}
}
}
impl Serializer for KeyRenderer {
type Ok = String;
type Error = WalkError;
type SerializeSeq = Impossible<String, WalkError>;
type SerializeTuple = Impossible<String, WalkError>;
type SerializeTupleStruct = Impossible<String, WalkError>;
type SerializeTupleVariant = Impossible<String, WalkError>;
type SerializeMap = Impossible<String, WalkError>;
type SerializeStruct = Impossible<String, WalkError>;
type SerializeStructVariant = Impossible<String, WalkError>;
fn serialize_str(self, value: &str) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_bool(self, value: bool) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_i64(self, value: i64) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_i128(self, value: i128) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_u64(self, value: u64) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_u128(self, value: u128) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_i8(self, value: i8) -> Result<String, WalkError> {
self.serialize_i64(i64::from(value))
}
fn serialize_i16(self, value: i16) -> Result<String, WalkError> {
self.serialize_i64(i64::from(value))
}
fn serialize_i32(self, value: i32) -> Result<String, WalkError> {
self.serialize_i64(i64::from(value))
}
fn serialize_u8(self, value: u8) -> Result<String, WalkError> {
self.serialize_u64(u64::from(value))
}
fn serialize_u16(self, value: u16) -> Result<String, WalkError> {
self.serialize_u64(u64::from(value))
}
fn serialize_u32(self, value: u32) -> Result<String, WalkError> {
self.serialize_u64(u64::from(value))
}
fn serialize_f32(self, value: f32) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_f64(self, value: f64) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_char(self, value: char) -> Result<String, WalkError> {
Ok(value.to_string())
}
fn serialize_bytes(self, value: &[u8]) -> Result<String, WalkError> {
Ok(format!("<{} bytes>", value.len()))
}
fn serialize_none(self) -> Result<String, WalkError> {
Ok("null".to_string())
}
fn serialize_some<T>(self, value: &T) -> Result<String, WalkError>
where
T: Serialize + ?Sized,
{
value.serialize(self)
}
fn serialize_unit(self) -> Result<String, WalkError> {
Ok("null".to_string())
}
fn serialize_unit_struct(self, name: &'static str) -> Result<String, WalkError> {
Ok(name.to_string())
}
fn serialize_unit_variant(
self,
name: &'static str,
variant_index: u32,
variant: &'static str,
) -> Result<String, WalkError> {
Ok(Self::variant_key(name, variant_index, variant))
}
fn serialize_newtype_struct<T>(self, name: &'static str, value: &T) -> Result<String, WalkError>
where
T: Serialize + ?Sized,
{
value
.serialize(self)
.map_err(|inner| WalkError::Custom(format!("inside newtype struct `{name}`: {inner}")))
}
fn serialize_newtype_variant<T>(
self,
name: &'static str,
variant_index: u32,
variant: &'static str,
value: &T,
) -> Result<String, WalkError>
where
T: Serialize + ?Sized,
{
let inner = value.serialize(KeyRenderer)?;
Ok(format!(
"{}({inner})",
Self::variant_key(name, variant_index, variant)
))
}
fn serialize_seq(self, len: Option<usize>) -> Result<Self::SerializeSeq, WalkError> {
Err(Self::not_a_scalar(match len {
Some(len) => format!("a sequence of {len} elements"),
None => "a sequence of unannounced length".to_string(),
}))
}
fn serialize_tuple(self, len: usize) -> Result<Self::SerializeTuple, WalkError> {
Err(Self::not_a_scalar(format_args!("a {len}-tuple")))
}
fn serialize_tuple_struct(
self,
name: &'static str,
len: usize,
) -> Result<Self::SerializeTupleStruct, WalkError> {
Err(Self::not_a_scalar(format_args!(
"tuple struct `{name}` with {len} fields"
)))
}
fn serialize_tuple_variant(
self,
name: &'static str,
variant_index: u32,
variant: &'static str,
len: usize,
) -> Result<Self::SerializeTupleVariant, WalkError> {
Err(Self::not_a_scalar(format_args!(
"tuple variant `{name}::{variant}` (variant #{variant_index}) with {len} fields"
)))
}
fn serialize_map(self, len: Option<usize>) -> Result<Self::SerializeMap, WalkError> {
Err(Self::not_a_scalar(match len {
Some(len) => format!("a map of {len} entries"),
None => "a map of unannounced length".to_string(),
}))
}
fn serialize_struct(
self,
name: &'static str,
len: usize,
) -> Result<Self::SerializeStruct, WalkError> {
Err(Self::not_a_scalar(format_args!(
"struct `{name}` with {len} fields"
)))
}
fn serialize_struct_variant(
self,
name: &'static str,
variant_index: u32,
variant: &'static str,
len: usize,
) -> Result<Self::SerializeStructVariant, WalkError> {
Err(Self::not_a_scalar(format_args!(
"struct variant `{name}::{variant}` (variant #{variant_index}) with {len} fields"
)))
}
}
impl<'a> Serializer for &'a mut FloatWalker {
type Ok = ();
type Error = WalkError;
type SerializeSeq = SeqWalker<'a>;
type SerializeTuple = SeqWalker<'a>;
type SerializeTupleStruct = SeqWalker<'a>;
type SerializeTupleVariant = VariantSeqWalker<'a>;
type SerializeMap = MapWalker<'a>;
type SerializeStruct = StructWalker<'a>;
type SerializeStructVariant = StructWalker<'a>;
fn serialize_f64(self, value: f64) -> Result<(), WalkError> {
self.check(value)
}
fn serialize_f32(self, value: f32) -> Result<(), WalkError> {
self.check(f64::from(value))
}
fn serialize_bool(self, _: bool) -> Result<(), WalkError> {
Ok(())
}
fn serialize_i8(self, value: i8) -> Result<(), WalkError> {
self.serialize_i64(i64::from(value))
}
fn serialize_i16(self, value: i16) -> Result<(), WalkError> {
self.serialize_i64(i64::from(value))
}
fn serialize_i32(self, value: i32) -> Result<(), WalkError> {
self.serialize_i64(i64::from(value))
}
fn serialize_i64(self, value: i64) -> Result<(), WalkError> {
self.serialize_i128(i128::from(value))
}
fn serialize_i128(self, _: i128) -> Result<(), WalkError> {
Ok(())
}
fn serialize_u8(self, value: u8) -> Result<(), WalkError> {
self.serialize_u64(u64::from(value))
}
fn serialize_u16(self, value: u16) -> Result<(), WalkError> {
self.serialize_u64(u64::from(value))
}
fn serialize_u32(self, value: u32) -> Result<(), WalkError> {
self.serialize_u64(u64::from(value))
}
fn serialize_u64(self, value: u64) -> Result<(), WalkError> {
self.serialize_u128(u128::from(value))
}
fn serialize_u128(self, _: u128) -> Result<(), WalkError> {
Ok(())
}
fn serialize_char(self, value: char) -> Result<(), WalkError> {
self.serialize_str(value.encode_utf8(&mut [0u8; 4]))
}
fn serialize_str(self, _: &str) -> Result<(), WalkError> {
Ok(())
}
fn serialize_bytes(self, _: &[u8]) -> Result<(), WalkError> {
Ok(())
}
fn serialize_none(self) -> Result<(), WalkError> {
Ok(())
}
fn serialize_some<T>(self, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
value.serialize(self)
}
fn serialize_unit(self) -> Result<(), WalkError> {
Ok(())
}
fn serialize_unit_struct(self, _: &'static str) -> Result<(), WalkError> {
Ok(())
}
fn serialize_unit_variant(
self,
_: &'static str,
_: u32,
_: &'static str,
) -> Result<(), WalkError> {
Ok(())
}
fn serialize_newtype_struct<T>(self, _: &'static str, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
value.serialize(self)
}
fn serialize_newtype_variant<T>(
self,
name: &'static str,
variant_index: u32,
variant: &'static str,
value: &T,
) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
assert!(
!variant.is_empty(),
"`{name}` variant #{variant_index} has an empty name; \
its path segment would be invisible"
);
let restore = self.push_field(variant);
let outcome = value.serialize(&mut *self);
self.pop_to(restore);
outcome
}
fn serialize_seq(self, len: Option<usize>) -> Result<SeqWalker<'a>, WalkError> {
Ok(SeqWalker::new(self, len, Origin::plain("a sequence")))
}
fn serialize_tuple(self, len: usize) -> Result<SeqWalker<'a>, WalkError> {
Ok(SeqWalker::new(self, Some(len), Origin::plain("a tuple")))
}
fn serialize_tuple_struct(
self,
name: &'static str,
len: usize,
) -> Result<SeqWalker<'a>, WalkError> {
Ok(SeqWalker::new(self, Some(len), Origin::plain(name)))
}
fn serialize_tuple_variant(
self,
name: &'static str,
variant_index: u32,
variant: &'static str,
len: usize,
) -> Result<VariantSeqWalker<'a>, WalkError> {
let origin = Origin::variant(name, variant_index, variant);
let restore = self.push_field(variant);
Ok(VariantSeqWalker {
seq: SeqWalker::new(self, Some(len), origin),
restore,
})
}
fn serialize_map(self, len: Option<usize>) -> Result<MapWalker<'a>, WalkError> {
Ok(MapWalker {
walker: self,
restore: None,
announced: len,
entries: 0,
origin: Origin::plain("a map"),
})
}
fn serialize_struct(
self,
name: &'static str,
len: usize,
) -> Result<StructWalker<'a>, WalkError> {
Ok(StructWalker {
walker: self,
restore: None,
announced: len,
fields: 0,
origin: Origin::plain(name),
})
}
fn serialize_struct_variant(
self,
name: &'static str,
variant_index: u32,
variant: &'static str,
len: usize,
) -> Result<StructWalker<'a>, WalkError> {
let origin = Origin::variant(name, variant_index, variant);
let restore = self.push_field(variant);
Ok(StructWalker {
walker: self,
restore: Some(restore),
announced: len,
fields: 0,
origin,
})
}
fn collect_str<T>(self, _: &T) -> Result<(), WalkError>
where
T: Display + ?Sized,
{
Ok(())
}
fn is_human_readable(&self) -> bool {
true
}
}
#[derive(Clone, Copy)]
struct Origin {
type_name: &'static str,
variant: Option<(&'static str, u32)>,
}
impl Origin {
fn plain(type_name: &'static str) -> Self {
Origin {
type_name,
variant: None,
}
}
fn variant(type_name: &'static str, variant_index: u32, variant: &'static str) -> Self {
Origin {
type_name,
variant: Some((variant, variant_index)),
}
}
}
impl Display for Origin {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.variant {
Some((variant, variant_index)) => write!(
f,
"`{}::{variant}` (variant #{variant_index})",
self.type_name
),
None => f.write_str(self.type_name),
}
}
}
fn assert_announced_len(origin: &Origin, announced: Option<usize>, emitted: usize) {
if let Some(announced) = announced {
assert_eq!(
emitted, announced,
"{origin} announced {announced} elements but emitted {emitted}"
);
}
}
struct SeqWalker<'a> {
walker: &'a mut FloatWalker,
index: usize,
announced: Option<usize>,
origin: Origin,
}
impl<'a> SeqWalker<'a> {
fn new(walker: &'a mut FloatWalker, announced: Option<usize>, origin: Origin) -> Self {
SeqWalker {
walker,
index: 0,
announced,
origin,
}
}
fn finish(&self) {
assert_announced_len(&self.origin, self.announced, self.index);
}
}
impl SerializeSeq for SeqWalker<'_> {
type Ok = ();
type Error = WalkError;
fn serialize_element<T>(&mut self, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
let restore = self.walker.push_index(self.index);
let outcome = value.serialize(&mut *self.walker);
self.walker.pop_to(restore);
self.index += 1;
outcome
}
fn end(self) -> Result<(), WalkError> {
self.finish();
Ok(())
}
}
impl SerializeTuple for SeqWalker<'_> {
type Ok = ();
type Error = WalkError;
fn serialize_element<T>(&mut self, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
SerializeSeq::serialize_element(self, value)
}
fn end(self) -> Result<(), WalkError> {
SerializeSeq::end(self)
}
}
impl SerializeTupleStruct for SeqWalker<'_> {
type Ok = ();
type Error = WalkError;
fn serialize_field<T>(&mut self, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
SerializeSeq::serialize_element(self, value)
}
fn end(self) -> Result<(), WalkError> {
SerializeSeq::end(self)
}
}
struct VariantSeqWalker<'a> {
seq: SeqWalker<'a>,
restore: usize,
}
impl SerializeTupleVariant for VariantSeqWalker<'_> {
type Ok = ();
type Error = WalkError;
fn serialize_field<T>(&mut self, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
SerializeSeq::serialize_element(&mut self.seq, value)
}
fn end(self) -> Result<(), WalkError> {
self.seq.finish();
self.seq.walker.pop_to(self.restore);
Ok(())
}
}
struct MapWalker<'a> {
walker: &'a mut FloatWalker,
restore: Option<usize>,
announced: Option<usize>,
entries: usize,
origin: Origin,
}
impl SerializeMap for MapWalker<'_> {
type Ok = ();
type Error = WalkError;
fn serialize_key<T>(&mut self, key: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
let rendered = key
.serialize(KeyRenderer)
.unwrap_or_else(|err| format!("<unrenderable key: {err}>"));
self.restore = Some(self.walker.push_field(&rendered));
self.entries += 1;
Ok(())
}
fn serialize_value<T>(&mut self, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
let outcome = value.serialize(&mut *self.walker);
if let Some(restore) = self.restore.take() {
self.walker.pop_to(restore);
}
outcome
}
fn end(self) -> Result<(), WalkError> {
assert_announced_len(&self.origin, self.announced, self.entries);
Ok(())
}
}
struct StructWalker<'a> {
walker: &'a mut FloatWalker,
restore: Option<usize>,
announced: usize,
fields: usize,
origin: Origin,
}
impl SerializeStruct for StructWalker<'_> {
type Ok = ();
type Error = WalkError;
fn serialize_field<T>(&mut self, key: &'static str, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
let restore = self.walker.push_field(key);
let outcome = value.serialize(&mut *self.walker);
self.walker.pop_to(restore);
self.fields += 1;
outcome
}
fn end(self) -> Result<(), WalkError> {
assert_announced_len(&self.origin, Some(self.announced), self.fields);
if let Some(restore) = self.restore {
self.walker.pop_to(restore);
}
Ok(())
}
}
impl SerializeStructVariant for StructWalker<'_> {
type Ok = ();
type Error = WalkError;
fn serialize_field<T>(&mut self, key: &'static str, value: &T) -> Result<(), WalkError>
where
T: Serialize + ?Sized,
{
SerializeStruct::serialize_field(self, key, value)
}
fn end(self) -> Result<(), WalkError> {
SerializeStruct::end(self)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde::Serialize;
use std::collections::BTreeMap;
#[derive(Serialize)]
struct Leaf {
edf: f64,
name: String,
}
#[derive(Serialize)]
struct Root {
blocks: Vec<Leaf>,
scale: Option<f64>,
counts: Vec<u32>,
by_term: BTreeMap<String, f64>,
}
fn root() -> Root {
Root {
blocks: vec![
Leaf {
edf: 1.0,
name: "a".to_string(),
},
Leaf {
edf: 2.0,
name: "b".to_string(),
},
],
scale: Some(0.5),
counts: vec![1, 2, 3],
by_term: BTreeMap::from([("s(x)".to_string(), 3.25)]),
}
}
#[test]
fn all_finite_payload_passes() {
assert!(ensure_serialized_floats_are_finite(&root()).is_ok());
}
#[test]
fn nested_sequence_element_reports_indexed_path() {
let mut value = root();
value.blocks[1].edf = f64::NAN;
let err = ensure_serialized_floats_are_finite(&value).unwrap_err();
assert_eq!(err.path, "blocks[1].edf");
assert!(err.value.is_nan());
assert!(
err.to_string().contains("blocks[1].edf must be finite"),
"message should name the path: {err}"
);
}
#[test]
fn optional_scalar_reports_its_field() {
let mut value = root();
value.scale = Some(f64::INFINITY);
let err = ensure_serialized_floats_are_finite(&value).unwrap_err();
assert_eq!(err.path, "scale");
assert_eq!(err.value, f64::INFINITY);
}
#[test]
fn map_value_reports_its_key() {
let mut value = root();
value
.by_term
.insert("s(z)".to_string(), f64::NEG_INFINITY);
let err = ensure_serialized_floats_are_finite(&value).unwrap_err();
assert_eq!(err.path, "by_term.s(z)");
}
#[test]
fn none_is_not_a_non_finite_float() {
let mut value = root();
value.scale = None;
assert!(ensure_serialized_floats_are_finite(&value).is_ok());
}
#[test]
fn bare_scalar_has_empty_path() {
let err = ensure_serialized_floats_are_finite(&f64::NAN).unwrap_err();
assert!(err.path.is_empty());
assert!(err.to_string().starts_with("value must be finite"));
}
#[test]
fn f32_non_finite_is_caught_and_widened() {
#[derive(Serialize)]
struct Small {
w: f32,
}
let err = ensure_serialized_floats_are_finite(&Small { w: f32::NAN }).unwrap_err();
assert_eq!(err.path, "w");
assert!(err.value.is_nan());
}
#[test]
fn struct_variant_and_newtype_variant_paths_are_reported() {
#[derive(Serialize)]
enum Node {
Scale { phi: f64 },
Raw(f64),
}
#[derive(Serialize)]
struct Holder {
node: Node,
}
let err =
ensure_serialized_floats_are_finite(&Holder { node: Node::Scale { phi: f64::NAN } })
.unwrap_err();
assert_eq!(err.path, "node.Scale.phi");
let err = ensure_serialized_floats_are_finite(&Holder {
node: Node::Raw(f64::INFINITY),
})
.unwrap_err();
assert_eq!(err.path, "node.Raw");
}
#[test]
fn path_state_is_restored_after_each_branch() {
#[derive(Serialize)]
struct Two {
first: Vec<Leaf>,
second: f64,
}
let err = ensure_serialized_floats_are_finite(&Two {
first: vec![Leaf {
edf: 1.0,
name: "ok".to_string(),
}],
second: f64::NAN,
})
.unwrap_err();
assert_eq!(err.path, "second");
}
#[test]
fn enum_keyed_map_reports_the_variant_as_its_key() {
#[derive(Serialize, PartialEq, Eq, PartialOrd, Ord)]
enum Term {
Linear,
Smooth,
}
#[derive(Serialize)]
struct Holder {
by_term: BTreeMap<Term, f64>,
}
let err = ensure_serialized_floats_are_finite(&Holder {
by_term: BTreeMap::from([(Term::Linear, 1.0), (Term::Smooth, f64::NAN)]),
})
.unwrap_err();
assert_eq!(err.path, "by_term.Smooth");
}
#[cfg(debug_assertions)]
#[test]
#[should_panic(expected = "announced 2 elements but emitted 1")]
fn a_compound_that_under_emits_its_announced_length_is_caught() {
struct UnderEmitting;
impl Serialize for UnderEmitting {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let mut seq = serializer.serialize_seq(Some(2))?;
seq.serialize_element(&1.0f64)?;
seq.end()
}
}
assert!(
ensure_serialized_floats_are_finite(&UnderEmitting).is_err(),
"a compound emitting fewer elements than it announced must not pass"
);
}
#[test]
fn walk_agrees_with_what_serde_json_would_write() {
let mut value = root();
value.blocks[0].edf = f64::NAN;
let json = serde_json::to_string(&value).expect("serde_json renders NaN as null");
assert!(
json.contains("\"edf\":null"),
"precondition: serde_json writes NaN as null, got {json}"
);
assert!(ensure_serialized_floats_are_finite(&value).is_err());
}
}