csm_core_lib/
hyperdim_serde.rs1use serde::de::{self, Visitor};
6use serde::{Deserialize, Deserializer, Serialize, Serializer};
7use std::fmt;
8
9use crate::hyperdim::HVec10240;
10
11impl Serialize for HVec10240 {
12 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
13 where
14 S: Serializer,
15 {
16 if serializer.is_human_readable() {
17 use base64::Engine;
19 use base64::engine::general_purpose::STANDARD;
20 let bytes = self.to_bytes();
21 let b64 = STANDARD.encode(&bytes);
22 serializer.serialize_str(&b64)
23 } else {
24 let bytes = self.to_bytes();
28 serializer.serialize_bytes(&bytes)
29 }
30 }
31}
32
33struct HVecVisitor;
34
35impl<'de> Visitor<'de> for HVecVisitor {
36 type Value = HVec10240;
37
38 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
39 formatter.write_str("a base64-encoded string, byte array, or a sequence of 80 u128 values")
40 }
41
42 fn visit_str<E>(self, v: &str) -> std::result::Result<Self::Value, E>
43 where
44 E: de::Error,
45 {
46 use base64::Engine;
47 use base64::engine::general_purpose::STANDARD;
48 let bytes = STANDARD.decode(v).map_err(de::Error::custom)?;
49 HVec10240::from_bytes(&bytes).map_err(de::Error::custom)
50 }
51
52 fn visit_bytes<E>(self, v: &[u8]) -> std::result::Result<Self::Value, E>
53 where
54 E: de::Error,
55 {
56 HVec10240::from_bytes(v).map_err(de::Error::custom)
57 }
58
59 fn visit_seq<A>(self, mut seq: A) -> std::result::Result<Self::Value, A::Error>
60 where
61 A: de::SeqAccess<'de>,
62 {
63 let mut words = Vec::with_capacity(80);
66 while let Some(word) = seq.next_element::<u128>()? {
67 words.push(word);
68 if words.len() > 1280 {
69 return Err(de::Error::custom("sequence too long for HVec10240"));
70 }
71 }
72
73 if words.len() == 80 {
74 let mut data = [0u128; 80];
75 data.copy_from_slice(&words);
76 Ok(HVec10240 { data })
77 } else if words.len() == 1280 {
78 #[allow(clippy::cast_possible_truncation)] let bytes: Vec<u8> = words.into_iter().map(|w| w as u8).collect();
81 HVec10240::from_bytes(&bytes).map_err(de::Error::custom)
82 } else {
83 Err(de::Error::custom(format!(
84 "expected 80 or 1280 elements, got {}",
85 words.len()
86 )))
87 }
88 }
89}
90
91impl<'de> Deserialize<'de> for HVec10240 {
92 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
93 where
94 D: Deserializer<'de>,
95 {
96 if deserializer.is_human_readable() {
101 deserializer.deserialize_any(HVecVisitor)
102 } else {
103 let bytes = <Vec<u8>>::deserialize(deserializer)?;
104 Self::from_bytes(&bytes).map_err(de::Error::custom)
105 }
106 }
107}