Skip to main content

yaml_rt_serde/value/
de.rs

1use std::fmt;
2
3use serde::de::{
4    self, DeserializeSeed, EnumAccess, IntoDeserializer, MapAccess, SeqAccess, VariantAccess,
5    Visitor,
6};
7use serde::{Deserialize, Deserializer};
8
9use super::{Mapping, Number, Tag, TaggedValue, Value};
10use crate::{Error, Result};
11
12impl<'de> Deserialize<'de> for Value {
13    fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
14    where
15        D: Deserializer<'de>,
16    {
17        deserializer.deserialize_any(ValueVisitor)
18    }
19}
20
21struct ValueVisitor;
22
23impl<'de> Visitor<'de> for ValueVisitor {
24    type Value = Value;
25
26    fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
27        formatter.write_str("any YAML value")
28    }
29
30    fn visit_unit<E>(self) -> std::result::Result<Value, E> {
31        Ok(Value::Null)
32    }
33
34    fn visit_none<E>(self) -> std::result::Result<Value, E> {
35        Ok(Value::Null)
36    }
37
38    fn visit_some<D>(self, deserializer: D) -> std::result::Result<Value, D::Error>
39    where
40        D: Deserializer<'de>,
41    {
42        Value::deserialize(deserializer)
43    }
44
45    fn visit_bool<E>(self, value: bool) -> std::result::Result<Value, E> {
46        Ok(Value::Bool(value))
47    }
48
49    fn visit_i8<E>(self, value: i8) -> std::result::Result<Value, E> {
50        Ok(Value::from(value))
51    }
52
53    fn visit_i16<E>(self, value: i16) -> std::result::Result<Value, E> {
54        Ok(Value::from(value))
55    }
56
57    fn visit_i32<E>(self, value: i32) -> std::result::Result<Value, E> {
58        Ok(Value::from(value))
59    }
60
61    fn visit_i64<E>(self, value: i64) -> std::result::Result<Value, E> {
62        Ok(Value::from(value))
63    }
64
65    fn visit_i128<E>(self, value: i128) -> std::result::Result<Value, E> {
66        Ok(Value::from(value))
67    }
68
69    fn visit_u8<E>(self, value: u8) -> std::result::Result<Value, E> {
70        Ok(Value::from(value))
71    }
72
73    fn visit_u16<E>(self, value: u16) -> std::result::Result<Value, E> {
74        Ok(Value::from(value))
75    }
76
77    fn visit_u32<E>(self, value: u32) -> std::result::Result<Value, E> {
78        Ok(Value::from(value))
79    }
80
81    fn visit_u64<E>(self, value: u64) -> std::result::Result<Value, E> {
82        Ok(Value::from(value))
83    }
84
85    fn visit_u128<E>(self, value: u128) -> std::result::Result<Value, E> {
86        Ok(Value::from(value))
87    }
88
89    fn visit_f32<E>(self, value: f32) -> std::result::Result<Value, E> {
90        Ok(Value::from(value))
91    }
92
93    fn visit_f64<E>(self, value: f64) -> std::result::Result<Value, E> {
94        Ok(Value::from(value))
95    }
96
97    fn visit_char<E>(self, value: char) -> std::result::Result<Value, E> {
98        Ok(Value::String(value.to_string()))
99    }
100
101    fn visit_str<E>(self, value: &str) -> std::result::Result<Value, E> {
102        Ok(Value::String(value.to_owned()))
103    }
104
105    fn visit_string<E>(self, value: String) -> std::result::Result<Value, E> {
106        Ok(Value::String(value))
107    }
108
109    fn visit_seq<A>(self, mut access: A) -> std::result::Result<Value, A::Error>
110    where
111        A: SeqAccess<'de>,
112    {
113        let mut values = Vec::with_capacity(access.size_hint().unwrap_or(0));
114        while let Some(value) = access.next_element()? {
115            values.push(value);
116        }
117        Ok(Value::Sequence(values))
118    }
119
120    fn visit_map<A>(self, mut access: A) -> std::result::Result<Value, A::Error>
121    where
122        A: MapAccess<'de>,
123    {
124        let mut mapping = Mapping::with_capacity(access.size_hint().unwrap_or(0));
125        while let Some((key, value)) = access.next_entry()? {
126            if mapping.insert(key, value).is_some() {
127                return Err(de::Error::custom("duplicate mapping key"));
128            }
129        }
130        Ok(Value::Mapping(mapping))
131    }
132
133    fn visit_enum<A>(self, access: A) -> std::result::Result<Value, A::Error>
134    where
135        A: EnumAccess<'de>,
136    {
137        let (tag, variant) = access.variant_seed(StringSeed)?;
138        let value = variant.newtype_variant::<Value>()?;
139        Ok(Value::Tagged(Box::new(TaggedValue {
140            tag: Tag::new(tag),
141            value,
142        })))
143    }
144}
145
146struct StringSeed;
147
148impl<'de> DeserializeSeed<'de> for StringSeed {
149    type Value = String;
150
151    fn deserialize<D>(self, deserializer: D) -> std::result::Result<String, D::Error>
152    where
153        D: Deserializer<'de>,
154    {
155        String::deserialize(deserializer)
156    }
157}
158
159macro_rules! deserialize_numeric_methods {
160    () => {
161        fn deserialize_i8<V>(self, visitor: V) -> Result<V::Value>
162        where
163            V: Visitor<'de>,
164        {
165            deserialize_signed(
166                self,
167                visitor,
168                |value| i8::try_from(value).ok(),
169                Visitor::visit_i8,
170            )
171        }
172
173        fn deserialize_i16<V>(self, visitor: V) -> Result<V::Value>
174        where
175            V: Visitor<'de>,
176        {
177            deserialize_signed(
178                self,
179                visitor,
180                |value| i16::try_from(value).ok(),
181                Visitor::visit_i16,
182            )
183        }
184
185        fn deserialize_i32<V>(self, visitor: V) -> Result<V::Value>
186        where
187            V: Visitor<'de>,
188        {
189            deserialize_signed(
190                self,
191                visitor,
192                |value| i32::try_from(value).ok(),
193                Visitor::visit_i32,
194            )
195        }
196
197        fn deserialize_i64<V>(self, visitor: V) -> Result<V::Value>
198        where
199            V: Visitor<'de>,
200        {
201            deserialize_signed(
202                self,
203                visitor,
204                |value| i64::try_from(value).ok(),
205                Visitor::visit_i64,
206            )
207        }
208
209        fn deserialize_i128<V>(self, visitor: V) -> Result<V::Value>
210        where
211            V: Visitor<'de>,
212        {
213            deserialize_signed(self, visitor, Some, Visitor::visit_i128)
214        }
215
216        fn deserialize_u8<V>(self, visitor: V) -> Result<V::Value>
217        where
218            V: Visitor<'de>,
219        {
220            deserialize_unsigned(
221                self,
222                visitor,
223                |value| u8::try_from(value).ok(),
224                Visitor::visit_u8,
225            )
226        }
227
228        fn deserialize_u16<V>(self, visitor: V) -> Result<V::Value>
229        where
230            V: Visitor<'de>,
231        {
232            deserialize_unsigned(
233                self,
234                visitor,
235                |value| u16::try_from(value).ok(),
236                Visitor::visit_u16,
237            )
238        }
239
240        fn deserialize_u32<V>(self, visitor: V) -> Result<V::Value>
241        where
242            V: Visitor<'de>,
243        {
244            deserialize_unsigned(
245                self,
246                visitor,
247                |value| u32::try_from(value).ok(),
248                Visitor::visit_u32,
249            )
250        }
251
252        fn deserialize_u64<V>(self, visitor: V) -> Result<V::Value>
253        where
254            V: Visitor<'de>,
255        {
256            deserialize_unsigned(
257                self,
258                visitor,
259                |value| u64::try_from(value).ok(),
260                Visitor::visit_u64,
261            )
262        }
263
264        fn deserialize_u128<V>(self, visitor: V) -> Result<V::Value>
265        where
266            V: Visitor<'de>,
267        {
268            deserialize_unsigned(self, visitor, Some, Visitor::visit_u128)
269        }
270
271        fn deserialize_f32<V>(self, visitor: V) -> Result<V::Value>
272        where
273            V: Visitor<'de>,
274        {
275            deserialize_float(self, visitor, |visitor, value| {
276                let value = checked_f64_to_f32(value)
277                    .ok_or_else(|| Error::message("expected an f32 in range"))?;
278                visitor.visit_f32(value)
279            })
280        }
281
282        fn deserialize_f64<V>(self, visitor: V) -> Result<V::Value>
283        where
284            V: Visitor<'de>,
285        {
286            deserialize_float(self, visitor, Visitor::visit_f64)
287        }
288    };
289}
290
291impl<'de> de::Deserializer<'de> for Value {
292    type Error = Error;
293
294    fn deserialize_any<V>(self, visitor: V) -> Result<V::Value>
295    where
296        V: Visitor<'de>,
297    {
298        match self {
299            Self::Null => visitor.visit_unit(),
300            Self::Bool(value) => visitor.visit_bool(value),
301            Self::Number(number) => visit_number(number, visitor),
302            Self::String(value) => visitor.visit_string(value),
303            Self::Sequence(values) => visitor.visit_seq(OwnedSeqAccess {
304                values: values.into_iter(),
305            }),
306            Self::Mapping(mapping) => visitor.visit_map(OwnedMapAccess {
307                entries: mapping.into_iter(),
308                pending: None,
309            }),
310            Self::Tagged(tagged) => visitor.visit_enum(OwnedTaggedAccess { tagged: *tagged }),
311        }
312    }
313
314    fn deserialize_option<V>(self, visitor: V) -> Result<V::Value>
315    where
316        V: Visitor<'de>,
317    {
318        if self.is_null() {
319            visitor.visit_none()
320        } else {
321            visitor.visit_some(self)
322        }
323    }
324
325    fn deserialize_enum<V>(
326        self,
327        _name: &'static str,
328        _variants: &'static [&'static str],
329        visitor: V,
330    ) -> Result<V::Value>
331    where
332        V: Visitor<'de>,
333    {
334        match self {
335            Self::Tagged(tagged) => visitor.visit_enum(OwnedTaggedAccess { tagged: *tagged }),
336            Self::String(value) => visitor.visit_enum(value.into_deserializer()),
337            _ => Err(Error::message("expected a YAML enum")),
338        }
339    }
340
341    fn deserialize_newtype_struct<V>(self, _name: &'static str, visitor: V) -> Result<V::Value>
342    where
343        V: Visitor<'de>,
344    {
345        visitor.visit_newtype_struct(self)
346    }
347
348    fn deserialize_ignored_any<V>(self, visitor: V) -> Result<V::Value>
349    where
350        V: Visitor<'de>,
351    {
352        visitor.visit_unit()
353    }
354
355    deserialize_numeric_methods!();
356
357    serde::forward_to_deserialize_any! {
358        bool char str string
359        bytes byte_buf unit unit_struct seq tuple tuple_struct map struct identifier
360    }
361
362    fn is_human_readable(&self) -> bool {
363        true
364    }
365}
366
367impl<'de> de::Deserializer<'de> for &'de Value {
368    type Error = Error;
369
370    fn deserialize_any<V>(self, visitor: V) -> Result<V::Value>
371    where
372        V: Visitor<'de>,
373    {
374        match self {
375            Value::Null => visitor.visit_unit(),
376            Value::Bool(value) => visitor.visit_bool(*value),
377            Value::Number(number) => visit_number(*number, visitor),
378            Value::String(value) => visitor.visit_borrowed_str(value),
379            Value::Sequence(values) => visitor.visit_seq(BorrowedSeqAccess {
380                values: values.iter(),
381            }),
382            Value::Mapping(mapping) => visitor.visit_map(BorrowedMapAccess {
383                entries: mapping.entries.iter(),
384                pending: None,
385            }),
386            Value::Tagged(tagged) => visitor.visit_enum(BorrowedTaggedAccess { tagged }),
387        }
388    }
389
390    fn deserialize_option<V>(self, visitor: V) -> Result<V::Value>
391    where
392        V: Visitor<'de>,
393    {
394        if self.is_null() {
395            visitor.visit_none()
396        } else {
397            visitor.visit_some(self)
398        }
399    }
400
401    fn deserialize_enum<V>(
402        self,
403        _name: &'static str,
404        _variants: &'static [&'static str],
405        visitor: V,
406    ) -> Result<V::Value>
407    where
408        V: Visitor<'de>,
409    {
410        match self {
411            Value::Tagged(tagged) => visitor.visit_enum(BorrowedTaggedAccess { tagged }),
412            Value::String(value) => visitor.visit_enum(
413                serde::de::value::BorrowedStrDeserializer::<Error>::new(value.as_str()),
414            ),
415            _ => Err(Error::message("expected a YAML enum")),
416        }
417    }
418
419    fn deserialize_newtype_struct<V>(self, _name: &'static str, visitor: V) -> Result<V::Value>
420    where
421        V: Visitor<'de>,
422    {
423        visitor.visit_newtype_struct(self)
424    }
425
426    fn deserialize_ignored_any<V>(self, visitor: V) -> Result<V::Value>
427    where
428        V: Visitor<'de>,
429    {
430        visitor.visit_unit()
431    }
432
433    deserialize_numeric_methods!();
434
435    serde::forward_to_deserialize_any! {
436        bool char str string
437        bytes byte_buf unit unit_struct seq tuple tuple_struct map struct identifier
438    }
439
440    fn is_human_readable(&self) -> bool {
441        true
442    }
443}
444
445fn visit_number<'de, V>(number: Number, visitor: V) -> Result<V::Value>
446where
447    V: Visitor<'de>,
448{
449    if number.is_f64() {
450        visitor.visit_f64(number.as_f64().expect("float is representable"))
451    } else if let Some(value) = number.as_i128()
452        && value < 0
453    {
454        visitor.visit_i128(value)
455    } else if let Some(value) = number.as_u128() {
456        visitor.visit_u128(value)
457    } else if let Some(value) = number.as_i128() {
458        visitor.visit_i128(value)
459    } else {
460        Err(Error::message("invalid YAML number"))
461    }
462}
463
464trait NumericValue {
465    fn into_number(self) -> Result<Number>;
466}
467
468impl NumericValue for Value {
469    fn into_number(self) -> Result<Number> {
470        match self {
471            Self::Number(number) => Ok(number),
472            _ => Err(Error::message("expected a number")),
473        }
474    }
475}
476
477impl NumericValue for &Value {
478    fn into_number(self) -> Result<Number> {
479        match self {
480            Value::Number(number) => Ok(*number),
481            _ => Err(Error::message("expected a number")),
482        }
483    }
484}
485
486fn deserialize_signed<'de, V, T>(
487    value: impl NumericValue,
488    visitor: V,
489    convert: impl FnOnce(i128) -> Option<T>,
490    visit: impl FnOnce(V, T) -> Result<V::Value>,
491) -> Result<V::Value>
492where
493    V: Visitor<'de>,
494{
495    let value = value
496        .into_number()?
497        .as_i128()
498        .and_then(convert)
499        .ok_or_else(|| Error::message("expected an integer in range"))?;
500    visit(visitor, value)
501}
502
503fn deserialize_unsigned<'de, V, T>(
504    value: impl NumericValue,
505    visitor: V,
506    convert: impl FnOnce(u128) -> Option<T>,
507    visit: impl FnOnce(V, T) -> Result<V::Value>,
508) -> Result<V::Value>
509where
510    V: Visitor<'de>,
511{
512    let value = value
513        .into_number()?
514        .as_u128()
515        .and_then(convert)
516        .ok_or_else(|| Error::message("expected an unsigned integer in range"))?;
517    visit(visitor, value)
518}
519
520fn deserialize_float<'de, V>(
521    value: impl NumericValue,
522    visitor: V,
523    visit: impl FnOnce(V, f64) -> Result<V::Value>,
524) -> Result<V::Value>
525where
526    V: Visitor<'de>,
527{
528    let value = value
529        .into_number()?
530        .as_f64()
531        .ok_or_else(|| Error::message("expected a number"))?;
532    visit(visitor, value)
533}
534
535fn checked_f64_to_f32(value: f64) -> Option<f32> {
536    #[expect(
537        clippy::cast_possible_truncation,
538        reason = "f32 deserialization applies Rust narrowing and then rejects finite overflow"
539    )]
540    let converted = value as f32;
541    (!value.is_finite() || converted.is_finite()).then_some(converted)
542}
543
544struct OwnedSeqAccess {
545    values: std::vec::IntoIter<Value>,
546}
547
548impl<'de> SeqAccess<'de> for OwnedSeqAccess {
549    type Error = Error;
550
551    fn next_element_seed<T>(&mut self, seed: T) -> Result<Option<T::Value>>
552    where
553        T: DeserializeSeed<'de>,
554    {
555        self.values
556            .next()
557            .map(|value| seed.deserialize(value))
558            .transpose()
559    }
560
561    fn size_hint(&self) -> Option<usize> {
562        Some(self.values.len())
563    }
564}
565
566struct BorrowedSeqAccess<'a> {
567    values: std::slice::Iter<'a, Value>,
568}
569
570impl<'de> SeqAccess<'de> for BorrowedSeqAccess<'de> {
571    type Error = Error;
572
573    fn next_element_seed<T>(&mut self, seed: T) -> Result<Option<T::Value>>
574    where
575        T: DeserializeSeed<'de>,
576    {
577        self.values
578            .next()
579            .map(|value| seed.deserialize(value))
580            .transpose()
581    }
582
583    fn size_hint(&self) -> Option<usize> {
584        Some(self.values.len())
585    }
586}
587
588struct OwnedMapAccess {
589    entries: super::IntoIter,
590    pending: Option<Value>,
591}
592
593impl<'de> MapAccess<'de> for OwnedMapAccess {
594    type Error = Error;
595
596    fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>>
597    where
598        K: DeserializeSeed<'de>,
599    {
600        let Some((key, value)) = self.entries.next() else {
601            return Ok(None);
602        };
603        self.pending = Some(value);
604        seed.deserialize(key).map(Some)
605    }
606
607    fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value>
608    where
609        V: DeserializeSeed<'de>,
610    {
611        seed.deserialize(
612            self.pending
613                .take()
614                .ok_or_else(|| Error::message("value requested without a key"))?,
615        )
616    }
617
618    fn size_hint(&self) -> Option<usize> {
619        Some(self.entries.len())
620    }
621}
622
623struct BorrowedMapAccess<'a> {
624    entries: std::slice::Iter<'a, (Value, Value)>,
625    pending: Option<&'a Value>,
626}
627
628impl<'de> MapAccess<'de> for BorrowedMapAccess<'de> {
629    type Error = Error;
630
631    fn next_key_seed<K>(&mut self, seed: K) -> Result<Option<K::Value>>
632    where
633        K: DeserializeSeed<'de>,
634    {
635        let Some((key, value)) = self.entries.next() else {
636            return Ok(None);
637        };
638        self.pending = Some(value);
639        seed.deserialize(key).map(Some)
640    }
641
642    fn next_value_seed<V>(&mut self, seed: V) -> Result<V::Value>
643    where
644        V: DeserializeSeed<'de>,
645    {
646        seed.deserialize(
647            self.pending
648                .take()
649                .ok_or_else(|| Error::message("value requested without a key"))?,
650        )
651    }
652
653    fn size_hint(&self) -> Option<usize> {
654        Some(self.entries.len())
655    }
656}
657
658struct OwnedTaggedAccess {
659    tagged: TaggedValue,
660}
661
662impl<'de> EnumAccess<'de> for OwnedTaggedAccess {
663    type Error = Error;
664    type Variant = OwnedTaggedVariant;
665
666    fn variant_seed<V>(self, seed: V) -> Result<(V::Value, Self::Variant)>
667    where
668        V: DeserializeSeed<'de>,
669    {
670        let variant = seed.deserialize(serde::de::value::StringDeserializer::<Error>::new(
671            self.tagged.tag.string,
672        ))?;
673        Ok((
674            variant,
675            OwnedTaggedVariant {
676                value: self.tagged.value,
677            },
678        ))
679    }
680}
681
682struct OwnedTaggedVariant {
683    value: Value,
684}
685
686impl<'de> VariantAccess<'de> for OwnedTaggedVariant {
687    type Error = Error;
688
689    fn unit_variant(self) -> Result<()> {
690        if self.value.is_null() {
691            Ok(())
692        } else {
693            Err(Error::message("expected a unit variant"))
694        }
695    }
696
697    fn newtype_variant_seed<T>(self, seed: T) -> Result<T::Value>
698    where
699        T: DeserializeSeed<'de>,
700    {
701        seed.deserialize(self.value)
702    }
703
704    fn tuple_variant<V>(self, _len: usize, visitor: V) -> Result<V::Value>
705    where
706        V: Visitor<'de>,
707    {
708        de::Deserializer::deserialize_seq(self.value, visitor)
709    }
710
711    fn struct_variant<V>(self, _fields: &'static [&'static str], visitor: V) -> Result<V::Value>
712    where
713        V: Visitor<'de>,
714    {
715        de::Deserializer::deserialize_map(self.value, visitor)
716    }
717}
718
719struct BorrowedTaggedAccess<'a> {
720    tagged: &'a TaggedValue,
721}
722
723impl<'de> EnumAccess<'de> for BorrowedTaggedAccess<'de> {
724    type Error = Error;
725    type Variant = BorrowedTaggedVariant<'de>;
726
727    fn variant_seed<V>(self, seed: V) -> Result<(V::Value, Self::Variant)>
728    where
729        V: DeserializeSeed<'de>,
730    {
731        let variant = seed.deserialize(serde::de::value::BorrowedStrDeserializer::<Error>::new(
732            self.tagged.tag.string.as_str(),
733        ))?;
734        Ok((
735            variant,
736            BorrowedTaggedVariant {
737                value: &self.tagged.value,
738            },
739        ))
740    }
741}
742
743struct BorrowedTaggedVariant<'a> {
744    value: &'a Value,
745}
746
747impl<'de> VariantAccess<'de> for BorrowedTaggedVariant<'de> {
748    type Error = Error;
749
750    fn unit_variant(self) -> Result<()> {
751        if self.value.is_null() {
752            Ok(())
753        } else {
754            Err(Error::message("expected a unit variant"))
755        }
756    }
757
758    fn newtype_variant_seed<T>(self, seed: T) -> Result<T::Value>
759    where
760        T: DeserializeSeed<'de>,
761    {
762        seed.deserialize(self.value)
763    }
764
765    fn tuple_variant<V>(self, _len: usize, visitor: V) -> Result<V::Value>
766    where
767        V: Visitor<'de>,
768    {
769        de::Deserializer::deserialize_seq(self.value, visitor)
770    }
771
772    fn struct_variant<V>(self, _fields: &'static [&'static str], visitor: V) -> Result<V::Value>
773    where
774        V: Visitor<'de>,
775    {
776        de::Deserializer::deserialize_map(self.value, visitor)
777    }
778}
779
780impl<'de> de::IntoDeserializer<'de, Error> for Value {
781    type Deserializer = Self;
782
783    fn into_deserializer(self) -> Self::Deserializer {
784        self
785    }
786}
787
788impl<'de> de::IntoDeserializer<'de, Error> for &'de Value {
789    type Deserializer = Self;
790
791    fn into_deserializer(self) -> Self::Deserializer {
792        self
793    }
794}