You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Rust中大型固定长度数组结构体的Serde序列化与反序列化问题

解决Serde序列化大固定长度数组的问题

由于Serde默认仅支持较小长度的固定数组序列化,对于[T; 4096]这类大数组,我们可以通过手动实现Serialize和Deserialize trait,将数组视为切片或原始字节序列处理,避免复制和中间Vec,同时不在非自描述格式中存储数组长度。

手动实现Serialize

序列化时直接将数组转为切片(或字节切片),利用Serde对切片的原生支持,无需复制数据:

use serde::ser::{Serialize, Serializer};

const COUNT: usize = 4096;

struct MyStruct {
    a: [u16; COUNT],
    b: [u8; COUNT]
}

impl Serialize for MyStruct {
    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
    where
        S: Serializer,
    {
        let mut struct_state = serializer.serialize_struct("MyStruct", 2)?;
        
        // 将u16数组转为字节切片(小端字节序,需与反序列化逻辑一致)
        let a_raw_bytes = unsafe {
            std::slice::from_raw_parts(
                self.a.as_ptr() as *const u8,
                COUNT * std::mem::size_of::<u16>(),
            )
        };
        struct_state.serialize_field("a", a_raw_bytes)?;
        
        // u8数组直接转为切片序列化
        struct_state.serialize_field("b", self.b.as_slice())?;
        
        struct_state.end()
    }
}

关键说明

  • as_slice() 创建数组的视图,无内存复制操作
  • 对于u16数组,转为字节切片后序列化,确保postcard等二进制格式仅写入原始字节,不附加长度信息
  • 字节序需与反序列化逻辑保持一致(示例中使用小端)

手动实现Deserialize

通过Visitor模式,直接读取固定长度的数据填充到数组,避免中间Vec:

use serde::de::{Deserialize, Deserializer, MapAccess, SeqAccess, Visitor};
use std::fmt;

impl<'de> Deserialize<'de> for MyStruct {
    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
    where
        D: Deserializer<'de>,
    {
        // 定义结构体字段的枚举标识
        enum Field { A, B }

        impl<'de> Deserialize<'de> for Field {
            fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
            where
                D: Deserializer<'de>,
            {
                struct FieldVisitor;

                impl<'de> Visitor<'de> for FieldVisitor {
                    type Value = Field;

                    fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
                        formatter.write_str("字段`a`或`b`")
                    }

                    fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
                    where
                        E: serde::de::Error,
                    {
                        match value {
                            "a" => Ok(Field::A),
                            "b" => Ok(Field::B),
                            _ => Err(serde::de::Error::unknown_field(value, &["a", "b"])),
                        }
                    }
                }

                deserializer.deserialize_identifier(FieldVisitor)
            }
        }

        // 定义MyStruct的反序列化Visitor
        struct MyStructVisitor;

        impl<'de> Visitor<'de> for MyStructVisitor {
            type Value = MyStruct;

            fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
                formatter.write_str("结构体MyStruct")
            }

            fn visit_map<V>(self, mut map: V) -> Result<Self::Value, V::Error>
            where
                V: MapAccess<'de>,
            {
                let mut a: Option<[u16; COUNT]> = None;
                let mut b: Option<[u8; COUNT]> = None;

                while let Some(key) = map.next_key()? {
                    match key {
                        Field::A => {
                            if a.is_some() {
                                return Err(serde::de::Error::duplicate_field("a"));
                            }
                            // 反序列化u16数组:读取固定长度字节并转换
                            a = Some(map.next_value_with(|deserializer| {
                                struct U16ArrayVisitor;

                                impl<'de> Visitor<'de> for U16ArrayVisitor {
                                    type Value = [u16; COUNT];

                                    fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
                                        write!(formatter, "包含{}个u16元素的数组", COUNT)
                                    }

                                    // 处理二进制格式(如postcard)的字节输入
                                    fn visit_bytes<E>(self, bytes: &[u8]) -> Result<Self::Value, E>
                                    where
                                        E: serde::de::Error,
                                    {
                                        if bytes.len() != COUNT * 2 {
                                            return Err(serde::de::Error::invalid_length(bytes.len(), &self));
                                        }
                                        let mut arr = [0; COUNT];
                                        // 将字节切片按小端转换为u16数组
                                        for (i, chunk) in bytes.chunks_exact(2).enumerate() {
                                            arr[i] = u16::from_le_bytes(chunk.try_into().unwrap());
                                        }
                                        Ok(arr)
                                    }

                                    // 处理序列格式(如JSON)的输入
                                    fn visit_seq<S>(self, mut seq: S) -> Result<Self::Value, S::Error>
                                    where
                                        S: SeqAccess<'de>,
                                    {
                                        let mut arr = [0; COUNT];
                                        for i in 0..COUNT {
                                            arr[i] = seq.next_element()?.ok_or_else(|| {
                                                serde::de::Error::invalid_length(i, &self)
                                            })?;
                                        }
                                        Ok(arr)
                                    }
                                }

                                deserializer.deserialize_bytes(U16ArrayVisitor)
                            })?)
                        }
                        Field::B => {
                            if b.is_some() {
                                return Err(serde::de::Error::duplicate_field("b"));
                            }
                            // 反序列化u8数组:直接读取固定长度字节
                            b = Some(map.next_value_with(|deserializer| {
                                struct U8ArrayVisitor;

                                impl<'de> Visitor<'de> for U8ArrayVisitor {
                                    type Value = [u8; COUNT];

                                    fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
                                        write!(formatter, "包含{}个u8元素的数组", COUNT)
                                    }

                                    fn visit_bytes<E>(self, bytes: &[u8]) -> Result<Self::Value, E>
                                    where
                                        E: serde::de::Error,
                                    {
                                        if bytes.len() != COUNT {
                                            return Err(serde::de::Error::invalid_length(bytes.len(), &self));
                                        }
                                        let mut arr = [0; COUNT];
                                        arr.copy_from_slice(bytes);
                                        Ok(arr)
                                    }

                                    fn visit_seq<S>(self, mut seq: S) -> Result<Self::Value, S::Error>
                                    where
                                        S: SeqAccess<'de>,
                                    {
                                        let mut arr = [0; COUNT];
                                        for i in 0..COUNT {
                                            arr[i] = seq.next_element()?.ok_or_else(|| {
                                                serde::de::Error::invalid_length(i, &self)
                                            })?;
                                        }
                                        Ok(arr)
                                    }
                                }

                                deserializer.deserialize_bytes(U8ArrayVisitor)
                            })?)
                        }
                    }
                }

                let a = a.ok_or_else(|| serde::de::Error::missing_field("a"))?;
                let b = b.ok_or_else(|| serde::de::Error::missing_field("b"))?;

                Ok(MyStruct { a, b })
            }
        }

        deserializer.deserialize_struct("MyStruct", &["a", "b"], MyStructVisitor)
    }
}

关键说明

  • Visitor模式允许精确控制反序列化过程,直接将输入数据填充到固定长度数组中
  • 针对二进制格式(如postcard),直接读取对应长度的字节并复制到数组(这是必要的内存拷贝,数据需从输入缓冲区转移到结构体数组,但无中间Vec开销)
  • 同时兼容JSON等序列格式的输入
  • 严格校验输入长度,确保符合COUNT的固定要求

验证postcard格式行为

使用postcard序列化时,会直接写入原始字节,不附加长度信息:

use postcard::{to_stdvec, from_bytes};

fn main() {
    let my_struct = MyStruct {
        a: [0x1234; COUNT],
        b: [0xAB; COUNT],
    };

    // 序列化
    let serialized = to_stdvec(&my_struct).unwrap();
    // 验证长度:COUNT*2(a的u16) + COUNT(b的u8) = COUNT*3
    assert_eq!(serialized.len(), COUNT * 3);

    // 反序列化
    let deserialized: MyStruct = from_bytes(&serialized).unwrap();
    assert_eq!(deserialized.a, my_struct.a);
    assert_eq!(deserialized.b, my_struct.b);
}

内容的提问来源于stack exchange,提问作者kiv_apple

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 21:54:55