use core::{
cell::RefCell,
convert::Infallible,
future::Future,
pin::{pin, Pin},
task::Context,
};
use crate::object_dict::{find_object, ODEntry};
use futures::{pending, task::noop_waker_ref};
use defmt_or_log::{debug, warn};
#[derive(Debug, Copy, Clone, PartialEq)]
#[repr(u8)]
pub enum NodeType {
ObjectValue = 1,
Unknown,
}
impl NodeType {
pub fn from_byte(b: u8) -> Self {
match b {
1 => Self::ObjectValue,
_ => Self::Unknown,
}
}
}
async fn write_bytes(bytes: &[u8], reg: &RefCell<u8>) {
for b in bytes {
*reg.borrow_mut() = *b;
pending!()
}
}
async fn serialize_object(obj: &ODEntry<'_>, sub: u8, reg: &RefCell<u8>) {
let data_size = obj.data.read_size(sub).unwrap() as u16;
let node_size = data_size + 4;
write_bytes(&node_size.to_le_bytes(), reg).await;
write_bytes(&[NodeType::ObjectValue as u8], reg).await;
write_bytes(&obj.index.to_le_bytes(), reg).await;
write_bytes(&[sub], reg).await;
const CHUNK_SIZE: usize = 32;
let mut buf = [0u8; CHUNK_SIZE];
let mut read_pos = 0;
loop {
obj.data.read(sub, read_pos, &mut buf).unwrap();
let copy_len = data_size as usize - read_pos;
read_pos += copy_len;
write_bytes(&buf[0..copy_len], reg).await;
if read_pos >= data_size as usize {
break;
}
}
}
async fn serialize_sm(objects: &[ODEntry<'_>], reg: &RefCell<u8>) {
for obj in objects {
let max_sub = obj.data.max_sub_number();
for sub in 0..max_sub + 1 {
let info = obj.data.sub_info(sub);
if info.is_err() {
continue;
}
let info = info.unwrap();
if !info.persist {
continue;
}
serialize_object(obj, sub, reg).await;
}
}
}
pub fn serialized_size(objects: &[ODEntry]) -> usize {
const OVERHEAD_SIZE: usize = 6;
let mut size = 0;
for obj in objects {
let max_sub = obj.data.max_sub_number();
for sub in 0..max_sub + 1 {
let info = obj.data.sub_info(sub);
if info.is_err() {
continue;
}
let info = info.unwrap();
if !info.persist {
continue;
}
let data_size = obj.data.read_size(sub).unwrap();
size += data_size + OVERHEAD_SIZE;
}
}
size
}
struct PersistSerializer<'a, 'b, F: Future> {
f: Pin<&'a mut F>,
reg: &'b RefCell<u8>,
}
impl<'a, 'b, F: Future> PersistSerializer<'a, 'b, F> {
pub fn new(f: Pin<&'a mut F>, reg: &'b RefCell<u8>) -> Self {
Self { f, reg }
}
}
impl<F: Future> embedded_io::ErrorType for PersistSerializer<'_, '_, F> {
type Error = Infallible;
}
impl<F: Future> embedded_io::Read for PersistSerializer<'_, '_, F> {
fn read(&mut self, buf: &mut [u8]) -> Result<usize, Infallible> {
let mut cx = Context::from_waker(noop_waker_ref());
let mut pos = 0;
loop {
if pos >= buf.len() {
return Ok(pos);
}
match self.f.as_mut().poll(&mut cx) {
core::task::Poll::Ready(_) => return Ok(pos),
core::task::Poll::Pending => {
buf[pos] = *self.reg.borrow();
pos += 1;
}
}
}
}
}
pub fn serialize(
od: &[ODEntry],
callback: &dyn Fn(&mut dyn embedded_io::Read<Error = Infallible>, usize),
) {
let reg = RefCell::new(0);
let fut = pin!(serialize_sm(od, ®));
let mut serializer = PersistSerializer::new(fut, ®);
let size = serialized_size(od);
callback(&mut serializer, size)
}
pub enum PersistReadError {
NodeLengthShort,
}
#[derive(Debug, PartialEq)]
pub struct ObjectValue<'a> {
pub index: u16,
pub sub: u8,
pub data: &'a [u8],
}
#[derive(Debug, PartialEq)]
pub enum PersistNodeRef<'a> {
ObjectValue(ObjectValue<'a>),
Unknown(&'a [u8]),
}
impl<'a> PersistNodeRef<'a> {
pub fn from_slice(data: &'a [u8]) -> Result<Self, PersistReadError> {
if data.is_empty() {
return Err(PersistReadError::NodeLengthShort);
}
match NodeType::from_byte(data[0]) {
NodeType::ObjectValue => {
if data.len() < 5 {
return Err(PersistReadError::NodeLengthShort);
}
Ok(Self::ObjectValue(ObjectValue {
index: u16::from_le_bytes(data[1..3].try_into().unwrap()),
sub: data[3],
data: &data[4..],
}))
}
NodeType::Unknown => Ok(PersistNodeRef::Unknown(data)),
}
}
}
struct PersistNodeReader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> PersistNodeReader<'a> {
pub fn new(data: &'a [u8]) -> Self {
Self { buf: data, pos: 0 }
}
}
impl<'a> Iterator for PersistNodeReader<'a> {
type Item = PersistNodeRef<'a>;
fn next(&mut self) -> Option<Self::Item> {
if self.buf.len() - self.pos < 2 {
return None;
}
let length = u16::from_le_bytes(self.buf[self.pos..self.pos + 2].try_into().unwrap());
self.pos += 2;
let node_slice = &self.buf[self.pos..self.pos + length as usize];
self.pos += length as usize;
PersistNodeRef::from_slice(node_slice).ok()
}
}
pub fn restore_stored_objects_ranged(
od: &[ODEntry],
stored_data: &[u8],
start_index: u16,
end_index: u16,
) {
let reader = PersistNodeReader::new(stored_data);
for item in reader {
match item {
PersistNodeRef::ObjectValue(restore) => {
if restore.index < start_index || restore.index > end_index {
continue;
}
if let Some(obj) = find_object(od, restore.index) {
if let Ok(_sub_info) = obj.sub_info(restore.sub) {
debug!(
"Restoring 0x{:x}sub{} with {:?}",
restore.index, restore.sub, restore.data
);
if let Err(abort_code) = obj.write(restore.sub, restore.data) {
warn!(
"Error restoring object 0x{:x}sub{}: {:x}",
restore.index, restore.sub, abort_code as u32
);
}
} else {
warn!(
"Saved object 0x{:x}sub{} not found in OD",
restore.index, restore.sub
);
}
} else {
warn!("Saved object 0x{:x} not found in OD", restore.index);
}
}
PersistNodeRef::Unknown(id) => warn!("Unknown persisted object read: {}", id[0]),
}
}
}
pub fn restore_stored_objects(od: &[ODEntry], stored_data: &[u8]) {
restore_stored_objects_ranged(od, stored_data, 0, u16::MAX);
}
pub fn restore_stored_comm_objects(od: &[ODEntry], stored_data: &[u8]) {
restore_stored_objects_ranged(od, stored_data, 0x1000, 0x1fff);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::object_dict::{
ConstField, NullTermByteField, ODEntry, ProvidesSubObjects, ScalarField, SubObjectAccess,
};
use zencan_common::objects::{DataType, ObjectCode, SubInfo};
use crate::persist::serialize;
#[test]
fn test_serialize_deserialize() {
#[derive(Default)]
struct Object100 {
value1: ScalarField<u32>,
value2: ScalarField<u16>,
}
impl ProvidesSubObjects for Object100 {
fn get_sub_object(&self, sub: u8) -> Option<(SubInfo, &dyn SubObjectAccess)> {
match sub {
0 => Some((
SubInfo::MAX_SUB_NUMBER,
const { &ConstField::new(2u8.to_le_bytes()) },
)),
1 => Some((
SubInfo {
size: 4,
data_type: DataType::UInt32,
persist: true,
..Default::default()
},
&self.value1,
)),
2 => Some((
SubInfo {
size: 4,
data_type: DataType::UInt32,
persist: false,
..Default::default()
},
&self.value2,
)),
_ => None,
}
}
fn object_code(&self) -> ObjectCode {
ObjectCode::Record
}
}
#[derive(Default)]
struct Object200 {
string: NullTermByteField<15>,
}
impl ProvidesSubObjects for Object200 {
fn get_sub_object(&self, sub: u8) -> Option<(SubInfo, &dyn SubObjectAccess)> {
match sub {
0 => Some((
SubInfo::new_visibile_str(self.string.len()).persist(true),
&self.string,
)),
_ => None,
}
}
fn object_code(&self) -> ObjectCode {
ObjectCode::Var
}
}
let inst100 = Box::leak(Box::new(Object100::default()));
let inst200 = Box::leak(Box::new(Object200::default()));
let od = Box::leak(Box::new([
ODEntry {
index: 0x100,
data: inst100,
},
ODEntry {
index: 0x200,
data: inst200,
},
]));
inst100.value1.store(42);
inst200.string.set_str("test".as_bytes()).unwrap();
let data = RefCell::new(Vec::new());
serialize(od, &|reader, _size| {
const CHUNK_SIZE: usize = 2;
let mut buf = [0; CHUNK_SIZE];
loop {
let n = reader.read(&mut buf).unwrap();
data.borrow_mut().extend_from_slice(&buf[..n]);
if n < buf.len() {
break;
}
}
});
let data = data.take();
assert_eq!(20, data.len());
assert_eq!(data.len(), serialized_size(od));
let mut deser = PersistNodeReader::new(&data);
assert_eq!(
deser.next().unwrap(),
PersistNodeRef::ObjectValue(ObjectValue {
index: 0x100,
sub: 1,
data: &42u32.to_le_bytes()
})
);
assert_eq!(
deser.next().unwrap(),
PersistNodeRef::ObjectValue(ObjectValue {
index: 0x200,
sub: 0,
data: "test".as_bytes()
})
);
assert_eq!(deser.next(), None);
}
}