Skip to main content

webdataset_core/
value.rs

1//! The dynamically typed values that make up a [`Sample`](crate::Sample).
2//!
3//! A WebDataset sample is a bag of files, and what a file decodes to depends on
4//! its extension: `.txt` becomes a string, `.cls` an integer, `.json` a tree,
5//! `.npy` a tensor, `.jpg` an image. [`Value`] is the union of those
6//! possibilities, plus [`Value::Custom`] as an escape hatch for decoders the
7//! library does not know about.
8
9use alloc::sync::Arc;
10use core::any::Any;
11use core::fmt;
12
13use crate::fields::Fields;
14use bytes::Bytes;
15
16use crate::error::{Error, Result};
17use crate::prelude::*;
18use crate::tensor::Tensor;
19
20/// A value carried by a sample field.
21#[derive(Clone)]
22#[non_exhaustive]
23pub enum Value {
24    /// Absent or JSON `null`.
25    Null,
26    /// A boolean.
27    Bool(bool),
28    /// An integer, e.g. the contents of a `.cls` file.
29    Int(i64),
30    /// A floating point number.
31    Float(f64),
32    /// Text, e.g. the contents of a `.txt` file.
33    Text(String),
34    /// Undecoded (or deliberately raw) file contents.
35    Bytes(Bytes),
36    /// An ordered sequence, e.g. a JSON array or a collated batch column.
37    List(Vec<Value>),
38    /// A string-keyed mapping, e.g. a JSON object or an `.npz` archive.
39    Map(Fields),
40    /// A dense numeric array, e.g. from `.npy` or `.ten`.
41    Tensor(Tensor),
42    /// A decoded image.
43    #[cfg(feature = "image")]
44    Image(Arc<image::DynamicImage>),
45    /// Anything else a user-supplied decoder produced.
46    Custom(Arc<dyn CustomValue>),
47}
48
49/// The trait objects accepted by [`Value::Custom`].
50pub trait CustomValue: Any + fmt::Debug + Send + Sync {
51    /// Upcast so callers can downcast to the concrete type.
52    fn as_any(&self) -> &dyn Any;
53}
54
55impl<T: Any + fmt::Debug + Send + Sync> CustomValue for T {
56    fn as_any(&self) -> &dyn Any {
57        self
58    }
59}
60
61impl Value {
62    /// Wrap an arbitrary value so it can live inside a sample.
63    pub fn custom<T: CustomValue>(value: T) -> Value {
64        Value::Custom(Arc::new(value))
65    }
66
67    /// A short type name, useful in error messages.
68    pub fn type_name(&self) -> &'static str {
69        match self {
70            Value::Null => "null",
71            Value::Bool(_) => "bool",
72            Value::Int(_) => "int",
73            Value::Float(_) => "float",
74            Value::Text(_) => "text",
75            Value::Bytes(_) => "bytes",
76            Value::List(_) => "list",
77            Value::Map(_) => "map",
78            Value::Tensor(_) => "tensor",
79            #[cfg(feature = "image")]
80            Value::Image(_) => "image",
81            Value::Custom(_) => "custom",
82        }
83    }
84
85    /// Borrow the raw bytes, if this value is undecoded.
86    pub fn as_bytes(&self) -> Option<&Bytes> {
87        match self {
88            Value::Bytes(b) => Some(b),
89            _ => None,
90        }
91    }
92
93    /// Borrow the text, if this value is a string.
94    pub fn as_str(&self) -> Option<&str> {
95        match self {
96            Value::Text(s) => Some(s),
97            _ => None,
98        }
99    }
100
101    /// Read this value as an integer, accepting `Int`, `Bool` and whole `Float`s.
102    pub fn as_i64(&self) -> Option<i64> {
103        match self {
104            Value::Int(i) => Some(*i),
105            Value::Bool(b) => Some(*b as i64),
106            // Whole floats convert; anything fractional or out of range does
107            // not. Written without `f64::fract` so this works without `std`.
108            Value::Float(f) if *f == (*f as i64) as f64 => Some(*f as i64),
109            _ => None,
110        }
111    }
112
113    /// Read this value as a float, accepting `Float` and `Int`.
114    pub fn as_f64(&self) -> Option<f64> {
115        match self {
116            Value::Float(f) => Some(*f),
117            Value::Int(i) => Some(*i as f64),
118            _ => None,
119        }
120    }
121
122    /// Borrow the list elements, if this value is a list.
123    pub fn as_list(&self) -> Option<&[Value]> {
124        match self {
125            Value::List(v) => Some(v),
126            _ => None,
127        }
128    }
129
130    /// Borrow the map, if this value is a map.
131    pub fn as_map(&self) -> Option<&Fields> {
132        match self {
133            Value::Map(m) => Some(m),
134            _ => None,
135        }
136    }
137
138    /// Borrow the tensor, if this value is one.
139    pub fn as_tensor(&self) -> Option<&Tensor> {
140        match self {
141            Value::Tensor(t) => Some(t),
142            _ => None,
143        }
144    }
145
146    /// Borrow the decoded image, if this value is one.
147    #[cfg(feature = "image")]
148    pub fn as_image(&self) -> Option<&image::DynamicImage> {
149        match self {
150            Value::Image(i) => Some(i),
151            _ => None,
152        }
153    }
154
155    /// Downcast a [`Value::Custom`] to a concrete type.
156    pub fn downcast_ref<T: Any>(&self) -> Option<&T> {
157        match self {
158            // Dereference through the `Arc` explicitly: the blanket impl below
159            // also covers `Arc<dyn CustomValue>`, so `c.as_any()` would resolve
160            // to the `Arc` itself rather than the value inside it.
161            Value::Custom(c) => (**c).as_any().downcast_ref::<T>(),
162            _ => None,
163        }
164    }
165
166    /// Like [`Value::as_bytes`] but produces a descriptive error.
167    pub fn expect_bytes(&self, key: &str) -> Result<&Bytes> {
168        self.as_bytes().ok_or_else(|| Error::value(format!("{key}: expected bytes, found {}", self.type_name())))
169    }
170
171    /// Convert to a `serde_json::Value` where a faithful mapping exists.
172    #[cfg(feature = "json")]
173    ///
174    /// Bytes are rejected rather than silently mangled; tensors become nested
175    /// arrays of numbers; images and custom values are unsupported.
176    pub fn to_json(&self) -> Result<serde_json::Value> {
177        use serde_json::Value as J;
178        Ok(match self {
179            Value::Null => J::Null,
180            Value::Bool(b) => J::Bool(*b),
181            Value::Int(i) => J::Number((*i).into()),
182            Value::Float(f) => serde_json::Number::from_f64(*f).map(J::Number).unwrap_or(J::Null),
183            Value::Text(s) => J::String(s.clone()),
184            Value::List(v) => J::Array(v.iter().map(|x| x.to_json()).collect::<Result<_>>()?),
185            Value::Map(m) => {
186                let mut o = serde_json::Map::new();
187                for (k, v) in m {
188                    o.insert(k.clone(), v.to_json()?);
189                }
190                J::Object(o)
191            }
192            Value::Tensor(t) => J::Array(
193                t.to_f64_vec()
194                    .into_iter()
195                    .map(|f| serde_json::Number::from_f64(f).map(J::Number).unwrap_or(J::Null))
196                    .collect(),
197            ),
198            other => {
199                return Err(Error::unsupported(format!("cannot represent {} as json", other.type_name())));
200            }
201        })
202    }
203}
204
205#[cfg(feature = "json")]
206impl From<serde_json::Value> for Value {
207    fn from(v: serde_json::Value) -> Value {
208        use serde_json::Value as J;
209        match v {
210            J::Null => Value::Null,
211            J::Bool(b) => Value::Bool(b),
212            J::Number(n) => {
213                if let Some(i) = n.as_i64() {
214                    Value::Int(i)
215                } else {
216                    Value::Float(n.as_f64().unwrap_or(f64::NAN))
217                }
218            }
219            J::String(s) => Value::Text(s),
220            J::Array(a) => Value::List(a.into_iter().map(Value::from).collect()),
221            J::Object(o) => Value::Map(o.into_iter().map(|(k, v)| (k, Value::from(v))).collect()),
222        }
223    }
224}
225
226macro_rules! from_impl {
227    ($($t:ty => $variant:expr),* $(,)?) => {
228        $(impl From<$t> for Value {
229            fn from(v: $t) -> Value {
230                #[allow(clippy::redundant_closure_call)]
231                ($variant)(v)
232            }
233        })*
234    };
235}
236
237from_impl! {
238    bool => Value::Bool,
239    i64 => Value::Int,
240    f64 => Value::Float,
241    String => Value::Text,
242    Bytes => Value::Bytes,
243    Tensor => Value::Tensor,
244    Vec<Value> => Value::List,
245    Fields => Value::Map,
246}
247
248impl From<&str> for Value {
249    fn from(v: &str) -> Value {
250        Value::Text(v.to_string())
251    }
252}
253
254impl From<Vec<u8>> for Value {
255    fn from(v: Vec<u8>) -> Value {
256        Value::Bytes(Bytes::from(v))
257    }
258}
259
260impl From<i32> for Value {
261    fn from(v: i32) -> Value {
262        Value::Int(v as i64)
263    }
264}
265
266impl From<usize> for Value {
267    fn from(v: usize) -> Value {
268        Value::Int(v as i64)
269    }
270}
271
272impl From<f32> for Value {
273    fn from(v: f32) -> Value {
274        Value::Float(v as f64)
275    }
276}
277
278impl fmt::Debug for Value {
279    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
280        match self {
281            Value::Null => f.write_str("Null"),
282            Value::Bool(b) => write!(f, "Bool({b})"),
283            Value::Int(i) => write!(f, "Int({i})"),
284            Value::Float(x) => write!(f, "Float({x})"),
285            Value::Text(s) => write!(f, "Text({:?})", Truncated(s)),
286            Value::Bytes(b) => write!(f, "Bytes({} bytes)", b.len()),
287            Value::List(v) => f.debug_tuple("List").field(v).finish(),
288            Value::Map(m) => f.debug_tuple("Map").field(m).finish(),
289            Value::Tensor(t) => write!(f, "Tensor({} {:?})", t.dtype().long_name(), t.shape()),
290            #[cfg(feature = "image")]
291            Value::Image(i) => {
292                use image::GenericImageView;
293                let (w, h) = i.dimensions();
294                write!(f, "Image({w}x{h})")
295            }
296            Value::Custom(c) => write!(f, "Custom({c:?})"),
297        }
298    }
299}
300
301struct Truncated<'a>(&'a str);
302
303impl fmt::Debug for Truncated<'_> {
304    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
305        if self.0.len() <= 60 {
306            write!(f, "{}", self.0)
307        } else {
308            let cut = self.0.char_indices().nth(57).map(|(i, _)| i).unwrap_or(self.0.len());
309            write!(f, "{}...", &self.0[..cut])
310        }
311    }
312}
313
314impl PartialEq for Value {
315    fn eq(&self, other: &Value) -> bool {
316        match (self, other) {
317            (Value::Null, Value::Null) => true,
318            (Value::Bool(a), Value::Bool(b)) => a == b,
319            (Value::Int(a), Value::Int(b)) => a == b,
320            (Value::Float(a), Value::Float(b)) => a == b,
321            (Value::Text(a), Value::Text(b)) => a == b,
322            (Value::Bytes(a), Value::Bytes(b)) => a == b,
323            (Value::List(a), Value::List(b)) => a == b,
324            (Value::Map(a), Value::Map(b)) => a == b,
325            (Value::Tensor(a), Value::Tensor(b)) => a == b,
326            _ => false,
327        }
328    }
329}
330
331#[cfg(test)]
332mod tests {
333    use super::*;
334
335    #[cfg(feature = "json")]
336    #[test]
337    fn converts_json_both_ways() {
338        let json: serde_json::Value = serde_json::from_str(r#"{"a": 1, "b": [true, null, "x"]}"#).unwrap();
339        let value = Value::from(json.clone());
340        assert_eq!(value.to_json().unwrap(), json);
341    }
342
343    #[test]
344    fn coerces_numbers() {
345        assert_eq!(Value::Int(3).as_f64(), Some(3.0));
346        assert_eq!(Value::Float(3.0).as_i64(), Some(3));
347        assert_eq!(Value::Float(3.5).as_i64(), None);
348    }
349
350    #[test]
351    fn round_trips_custom_values() {
352        #[derive(Debug, PartialEq)]
353        struct Mine(u32);
354        let v = Value::custom(Mine(7));
355        assert_eq!(v.downcast_ref::<Mine>(), Some(&Mine(7)));
356        assert_eq!(v.downcast_ref::<u32>(), None);
357    }
358
359    #[test]
360    fn truncates_long_text_in_debug() {
361        let long = "x".repeat(200);
362        let shown = format!("{:?}", Value::Text(long));
363        assert!(shown.len() < 80, "{shown}");
364    }
365}