use core::marker::PhantomData;
use indexmap::IndexMap;
use serde::Deserialize;
use serde::Deserializer;
use serde::Serialize;
use serde::Serializer;
use serde::de::DeserializeSeed;
use serde::de::MapAccess;
use serde::de::Visitor;
use serde::ser::Error;
use serde::ser::SerializeMap;
struct IndexMapVisitor<V>(PhantomData<fn() -> IndexMap<Vec<u8>, V, ahash::RandomState>>);
impl<'de, V: Deserialize<'de>> Visitor<'de> for IndexMapVisitor<V> {
type Value = IndexMap<Vec<u8>, V, ahash::RandomState>;
fn expecting(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
formatter.write_str("map")
}
fn visit_map<M: MapAccess<'de>>(self, mut access: M) -> Result<Self::Value, M::Error> {
let mut map = IndexMap::with_capacity_and_hasher(
access.size_hint().unwrap_or(0),
ahash::RandomState::default(),
);
while let Some((k, v)) = access.next_entry_seed(ByteKey, PhantomData)? {
map.insert(k, v);
}
Ok(map)
}
}
struct ByteKey;
impl<'de> DeserializeSeed<'de> for ByteKey {
type Value = Vec<u8>;
fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<Self::Value, D::Error> {
deserializer.deserialize_str(self)
}
}
impl Visitor<'_> for ByteKey {
type Value = Vec<u8>;
fn expecting(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
formatter.write_str("string")
}
fn visit_str<E: serde::de::Error>(self, string: &str) -> Result<Self::Value, E> {
Ok(string.as_bytes().to_vec())
}
}
pub(crate) fn serialize<S: Serializer, T: Serialize>(
data: &IndexMap<Vec<u8>, T, ahash::RandomState>,
serializer: S,
) -> Result<S::Ok, S::Error> {
let mut map = serializer.serialize_map(Some(data.len()))?;
for (k, v) in data {
match core::str::from_utf8(k) {
Ok(k) => map.serialize_entry(k, &v)?,
Err(e) => return Err(Error::custom(e)),
}
}
map.end()
}
pub(crate) fn deserialize<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
deserializer: D,
) -> Result<IndexMap<Vec<u8>, T, ahash::RandomState>, D::Error> {
deserializer.deserialize_map(IndexMapVisitor(PhantomData))
}
#[cfg(test)]
mod test {
use indexmap::IndexMap;
use serde::Deserialize;
use serde::Serialize;
#[derive(Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
struct Map {
#[serde(with = "super")]
data: IndexMap<Vec<u8>, usize, ahash::RandomState>,
}
#[test]
fn test() {
let mut map = Map::default();
map.data.insert(b"a".to_vec(), 1);
map.data.insert("ツ".as_bytes().to_vec(), 2);
assert_eq!(round_trip_json_str(&map), map);
assert_eq!(round_trip_json_vec(&map), map);
}
#[track_caller]
fn round_trip_json_str(map: &Map) -> Map {
serde_json::from_str(&serde_json::to_string(map).unwrap()).unwrap()
}
#[track_caller]
fn round_trip_json_vec(map: &Map) -> Map {
serde_json::de::from_slice(&serde_json::to_vec(map).unwrap()).unwrap()
}
}