use core::cell::UnsafeCell;
use zencan_common::{
i24, sdo::AbortCode, traits::ReadSize, u24, AtomicCell, TimeDifference, TimeOfDay,
};
pub trait SubObjectAccess: Sync + Send {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode>;
fn read_size(&self) -> usize;
fn write(&self, data: &[u8]) -> Result<(), AbortCode>;
fn begin_partial(&self) -> Result<(), AbortCode> {
Err(AbortCode::UnsupportedAccess)
}
fn write_partial(&self, _buf: &[u8]) -> Result<(), AbortCode> {
Err(AbortCode::UnsupportedAccess)
}
fn end_partial(&self) -> Result<(), AbortCode> {
Err(AbortCode::UnsupportedAccess)
}
}
#[allow(missing_debug_implementations)]
pub struct ScalarField<T: Copy> {
value: AtomicCell<T>,
}
impl<T: Send + Copy + PartialEq> ScalarField<T> {
pub fn load(&self) -> T {
self.value.load()
}
pub fn store(&self, value: T) {
self.value.store(value);
}
}
impl<T: Copy + Default> Default for ScalarField<T> {
fn default() -> Self {
Self {
value: AtomicCell::default(),
}
}
}
macro_rules! impl_scalar_field {
($rust_type: ty) => {
impl ScalarField<$rust_type> {
pub const fn new(value: $rust_type) -> Self {
Self {
value: AtomicCell::new(value),
}
}
}
impl SubObjectAccess for ScalarField<$rust_type> {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
let bytes = self.value.load().to_le_bytes();
if offset < bytes.len() {
let read_len = buf.len().min(bytes.len() - offset);
buf[0..read_len].copy_from_slice(&bytes[offset..offset + read_len]);
Ok(read_len)
} else {
Ok(0)
}
}
fn read_size(&self) -> usize {
<$rust_type as ReadSize>::READ_SIZE
}
fn write(&self, data: &[u8]) -> Result<(), AbortCode> {
let value = <$rust_type>::from_le_bytes(data.try_into().map_err(|_| {
if data.len() < size_of::<$rust_type>() {
AbortCode::DataTypeMismatchLengthLow
} else {
AbortCode::DataTypeMismatchLengthHigh
}
})?);
self.value.store(value);
Ok(())
}
}
};
}
impl_scalar_field!(u8);
impl_scalar_field!(u16);
impl_scalar_field!(u24);
impl_scalar_field!(u32);
impl_scalar_field!(u64);
impl_scalar_field!(i8);
impl_scalar_field!(i16);
impl_scalar_field!(i24);
impl_scalar_field!(i32);
impl_scalar_field!(i64);
impl_scalar_field!(f32);
impl_scalar_field!(f64);
impl SubObjectAccess for ScalarField<bool> {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
let value = self.value.load();
if offset != 0 || buf.len() > 1 {
return Err(AbortCode::DataTypeMismatchLengthHigh);
}
buf[0] = if value { 1 } else { 0 };
Ok(1)
}
fn read_size(&self) -> usize {
1
}
fn write(&self, data: &[u8]) -> Result<(), AbortCode> {
if data.len() != 1 {
return Err(AbortCode::DataTypeMismatchLengthHigh);
}
let value = data[0] != 0;
self.value.store(value);
Ok(())
}
}
impl SubObjectAccess for ScalarField<TimeDifference> {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
let value = self.value.load();
let bytes = value.to_le_bytes();
if offset < bytes.len() {
let read_len = buf.len().min(bytes.len() - offset);
buf[0..read_len].copy_from_slice(&bytes[offset..offset + read_len]);
Ok(read_len)
} else {
Ok(0)
}
}
fn read_size(&self) -> usize {
6
}
fn write(&self, data: &[u8]) -> Result<(), AbortCode> {
let value = TimeDifference::from_le_bytes(data.try_into().map_err(|_| {
if data.len() < 6 {
AbortCode::DataTypeMismatchLengthLow
} else {
AbortCode::DataTypeMismatchLengthHigh
}
})?);
self.value.store(value);
Ok(())
}
}
impl ScalarField<TimeDifference> {
pub const fn new(value: TimeDifference) -> Self {
Self {
value: AtomicCell::new(value),
}
}
}
impl SubObjectAccess for ScalarField<TimeOfDay> {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
let value = self.value.load();
let bytes = value.to_le_bytes();
if offset < bytes.len() {
let read_len = buf.len().min(bytes.len() - offset);
buf[0..read_len].copy_from_slice(&bytes[offset..offset + read_len]);
Ok(read_len)
} else {
Ok(0)
}
}
fn read_size(&self) -> usize {
6
}
fn write(&self, data: &[u8]) -> Result<(), AbortCode> {
let value = TimeOfDay::from_le_bytes(data.try_into().map_err(|_| {
if data.len() < 6 {
AbortCode::DataTypeMismatchLengthLow
} else {
AbortCode::DataTypeMismatchLengthHigh
}
})?);
self.value.store(value);
Ok(())
}
}
impl ScalarField<TimeOfDay> {
pub const fn new(value: TimeOfDay) -> Self {
Self {
value: AtomicCell::new(value),
}
}
}
#[allow(clippy::len_without_is_empty, missing_debug_implementations)]
pub struct ByteField<const N: usize> {
value: UnsafeCell<[u8; N]>,
write_offset: AtomicCell<Option<usize>>,
}
unsafe impl<const N: usize> Sync for ByteField<N> {}
impl<const N: usize> ByteField<N> {
pub const fn new(value: [u8; N]) -> Self {
Self {
value: UnsafeCell::new(value),
write_offset: AtomicCell::new(None),
}
}
pub fn len(&self) -> usize {
N
}
pub fn store(&self, value: [u8; N]) {
self.write_offset.store(None);
critical_section::with(|_| {
let bytes = unsafe { &mut *self.value.get() };
bytes.copy_from_slice(&value);
});
}
pub fn load(&self) -> [u8; N] {
critical_section::with(|_| unsafe { *self.value.get() })
}
}
impl<const N: usize> Default for ByteField<N> {
fn default() -> Self {
Self {
value: UnsafeCell::new([0; N]),
write_offset: AtomicCell::new(None),
}
}
}
impl<const N: usize> SubObjectAccess for ByteField<N> {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
critical_section::with(|_| {
let bytes = unsafe { &*self.value.get() };
if bytes.len() > offset {
let read_len = buf.len().min(bytes.len() - offset);
buf[..read_len].copy_from_slice(&bytes[offset..offset + read_len]);
Ok(read_len)
} else {
Ok(0)
}
})
}
fn read_size(&self) -> usize {
N
}
fn write(&self, data: &[u8]) -> Result<(), AbortCode> {
critical_section::with(|_| {
let bytes = unsafe { &mut *self.value.get() };
if data.len() > bytes.len() {
return Err(AbortCode::DataTypeMismatchLengthHigh);
}
bytes[..data.len()].copy_from_slice(data);
Ok(())
})
}
fn begin_partial(&self) -> Result<(), AbortCode> {
self.write_offset.store(Some(0));
Ok(())
}
fn write_partial(&self, buf: &[u8]) -> Result<(), AbortCode> {
let offset = self
.write_offset
.fetch_update(|old| Some(old.map(|x| x + buf.len())))
.unwrap();
if offset.is_none() {
return Err(AbortCode::GeneralError);
}
let offset = offset.unwrap();
if offset + buf.len() > N {
return Err(AbortCode::DataTypeMismatchLengthHigh);
}
critical_section::with(|_| {
let bytes = unsafe { &mut *self.value.get() };
bytes[offset..offset + buf.len()].copy_from_slice(buf);
});
Ok(())
}
fn end_partial(&self) -> Result<(), AbortCode> {
self.write_offset.store(None);
Ok(())
}
}
#[allow(clippy::len_without_is_empty, missing_debug_implementations)]
pub struct NullTermByteField<const N: usize>(ByteField<N>);
impl<const N: usize> NullTermByteField<N> {
pub const fn new(value: [u8; N]) -> Self {
Self(ByteField::new(value))
}
pub fn len(&self) -> usize {
N
}
pub fn load(&self) -> [u8; N] {
self.0.load()
}
pub fn store(&self, value: [u8; N]) {
self.0.store(value);
}
pub fn set_str(&self, value: &[u8]) -> Result<(), AbortCode> {
self.0.begin_partial()?;
self.0.write_partial(value)?;
if value.len() < N {
self.0.write_partial(&[0])?;
}
self.end_partial()?;
Ok(())
}
}
impl<const N: usize> Default for NullTermByteField<N> {
fn default() -> Self {
Self(ByteField::default())
}
}
impl<const N: usize> SubObjectAccess for NullTermByteField<N> {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
let size = self.0.read(offset, buf)?;
let size = buf[0..size].iter().position(|b| *b == 0).unwrap_or(size);
Ok(size)
}
fn read_size(&self) -> usize {
critical_section::with(|_| {
let bytes = unsafe { &*self.0.value.get() };
bytes.iter().position(|b| *b == 0).unwrap_or(bytes.len())
})
}
fn write(&self, data: &[u8]) -> Result<(), AbortCode> {
self.0.begin_partial()?;
self.0.write_partial(data)?;
if data.len() < N {
self.0.write_partial(&[0])?;
}
self.0.end_partial()?;
Ok(())
}
fn begin_partial(&self) -> Result<(), AbortCode> {
self.0.begin_partial()
}
fn write_partial(&self, data: &[u8]) -> Result<(), AbortCode> {
self.0.write_partial(data)
}
fn end_partial(&self) -> Result<(), AbortCode> {
if self.0.write_offset.load().unwrap_or(0) < N {
self.0.write_partial(&[0])?;
}
self.0.end_partial()
}
}
#[derive(Clone, Copy, Debug)]
pub struct ConstByteRefField {
value: &'static [u8],
}
impl ConstByteRefField {
pub const fn new(value: &'static [u8]) -> Self {
Self { value }
}
}
impl SubObjectAccess for ConstByteRefField {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
let read_len = buf.len().min(self.value.len() - offset);
buf[..read_len].copy_from_slice(&self.value[offset..offset + read_len]);
Ok(read_len)
}
fn read_size(&self) -> usize {
self.value.len()
}
fn write(&self, _data: &[u8]) -> Result<(), AbortCode> {
Err(AbortCode::ReadOnly)
}
}
#[derive(Debug)]
pub struct ConstField<const N: usize> {
bytes: [u8; N],
}
impl<const N: usize> ConstField<N> {
pub const fn new(bytes: [u8; N]) -> Self {
Self { bytes }
}
}
impl<const N: usize> SubObjectAccess for ConstField<N> {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
if offset < self.bytes.len() {
let read_len = buf.len().min(self.bytes.len() - offset);
buf[..read_len].copy_from_slice(&self.bytes[offset..offset + read_len]);
Ok(read_len)
} else {
Ok(0)
}
}
fn read_size(&self) -> usize {
N
}
fn write(&self, _data: &[u8]) -> Result<(), AbortCode> {
Err(AbortCode::ReadOnly)
}
}
#[allow(missing_debug_implementations)]
pub struct CallbackSubObject {
handler: AtomicCell<Option<&'static dyn SubObjectAccess>>,
}
impl Default for CallbackSubObject {
fn default() -> Self {
Self::new()
}
}
impl CallbackSubObject {
pub const fn new() -> Self {
Self {
handler: AtomicCell::new(None),
}
}
pub fn register_handler(&self, handler: &'static dyn SubObjectAccess) {
self.handler.store(Some(handler));
}
}
impl SubObjectAccess for CallbackSubObject {
fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize, AbortCode> {
if let Some(handler) = self.handler.load() {
handler.read(offset, buf)
} else {
Err(AbortCode::ResourceNotAvailable)
}
}
fn read_size(&self) -> usize {
if let Some(handler) = self.handler.load() {
handler.read_size()
} else {
0
}
}
fn write(&self, data: &[u8]) -> Result<(), AbortCode> {
if let Some(handler) = self.handler.load() {
handler.write(data)
} else {
Err(AbortCode::ResourceNotAvailable)
}
}
fn begin_partial(&self) -> Result<(), AbortCode> {
if let Some(handler) = self.handler.load() {
handler.begin_partial()
} else {
Err(AbortCode::ResourceNotAvailable)
}
}
fn write_partial(&self, buf: &[u8]) -> Result<(), AbortCode> {
if let Some(handler) = self.handler.load() {
handler.write_partial(buf)
} else {
Err(AbortCode::ResourceNotAvailable)
}
}
fn end_partial(&self) -> Result<(), AbortCode> {
if let Some(handler) = self.handler.load() {
handler.end_partial()
} else {
Err(AbortCode::ResourceNotAvailable)
}
}
}
#[cfg(test)]
mod tests {
use zencan_common::objects::{ObjectCode, SubInfo};
use crate::object_dict::{ObjectAccess, ProvidesSubObjects};
use super::*;
#[derive(Default)]
struct ExampleRecord {
val1: ScalarField<u32>,
val2: ScalarField<bool>,
val3: NullTermByteField<10>,
}
impl ProvidesSubObjects for ExampleRecord {
fn get_sub_object(&self, sub: u8) -> Option<(SubInfo, &dyn SubObjectAccess)> {
match sub {
0 => Some((
SubInfo::MAX_SUB_NUMBER,
const { &ConstField::new(3u8.to_le_bytes()) },
)),
1 => Some((SubInfo::new_u32().rw_access(), &self.val1)),
2 => Some((SubInfo::new_u8().rw_access(), &self.val2)),
3 => Some((
SubInfo::new_visibile_str(self.val3.len()).rw_access(),
&self.val3,
)),
_ => None,
}
}
fn object_code(&self) -> ObjectCode {
ObjectCode::Record
}
}
#[test]
fn test_record_with_provides_sub_objects() {
let record = ExampleRecord::default();
assert_eq!(3, record.read_u8(0).unwrap());
record.write(1, &42u32.to_le_bytes()).unwrap();
assert_eq!(42, record.read_u32(1).unwrap());
record.begin_partial(3).unwrap();
record
.write_partial(3, &[0, 1, 2, 3, 4, 5, 6, 7, 8, 9])
.unwrap();
let mut buf = [0; 10];
record.read(3, 0, &mut buf).unwrap();
assert_eq!([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], buf);
record.begin_partial(3).unwrap();
record.write_partial(3, &[0, 1, 2, 3]).unwrap();
record.write_partial(3, &[4, 5, 6, 7]).unwrap();
record.end_partial(3).unwrap();
let mut buf = [0; 9];
record.read(3, 0, &mut buf).unwrap();
assert_eq!([0u8, 1, 2, 3, 4, 5, 6, 7, 0], buf)
}
fn sub_read_test_helper(field: &dyn SubObjectAccess, expected_bytes: &[u8]) {
let n = expected_bytes.len();
assert_eq!(n, field.read_size());
let mut read_buf = vec![0xffu8; n + 10];
let read_size = field.read(0, &mut read_buf).unwrap();
assert_eq!(n, read_size);
assert_eq!(expected_bytes, &read_buf[0..n]);
let mut read_buf = vec![0xffu8; n + 10];
let read_size = field.read(0, &mut read_buf).unwrap();
assert_eq!(n, read_size);
assert_eq!(expected_bytes, &read_buf[0..n]);
if n > 2 {
let mut read_buf = vec![0xffu8; n + 10];
let read_size = field.read(2, &mut read_buf).unwrap();
assert_eq!(n - 2, read_size);
assert_eq!(&expected_bytes[2..], &read_buf[0..n - 2]);
let mut read_buf = vec![0xffu8; n - 2];
let read_size = field.read(1, &mut read_buf).unwrap();
assert_eq!(n - 2, read_size);
assert_eq!(expected_bytes[1..n - 1], read_buf);
} else {
let mut read_buf = vec![0xffu8; n + 10];
let read_size = field.read(1, &mut read_buf).unwrap();
assert_eq!(n.saturating_sub(1), read_size);
assert_eq!(&expected_bytes[1..], &read_buf[0..read_size]);
}
}
#[test]
fn test_scalar_field_bool() {
let field = ScalarField::<bool>::default();
field.store(true);
assert_eq!(1, field.read_size());
let mut read_buf = [0xffu8; 1];
let read_size = field.read(0, &mut read_buf).unwrap();
assert_eq!(1, read_size);
assert_eq!([1], read_buf);
}
#[test]
fn test_scalar_field_u8() {
let field = ScalarField::<u8>::new(42u8);
let exp_bytes = 42u8.to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_u32() {
let field = ScalarField::<u32>::new(42u32);
let exp_bytes = 42u32.to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_u16() {
let field = ScalarField::<u16>::new(42u16);
let exp_bytes = 42u16.to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_u24() {
let field = ScalarField::<u24>::new(u24::new(42));
let exp_bytes = u24::new(42).to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_u64() {
let field = ScalarField::<u64>::new(42u64);
let exp_bytes = 42u64.to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_i8() {
let field = ScalarField::<i8>::new(-42i8);
let exp_bytes = (-42i8).to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_i16() {
let field = ScalarField::<i16>::new(-42i16);
let exp_bytes = (-42i16).to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_i24() {
let field = ScalarField::<i24>::new(i24::new(-42));
let exp_bytes = i24::new(-42).to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_i32() {
let field = ScalarField::<i32>::new(-42i32);
let exp_bytes = (-42i32).to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_i64() {
let field = ScalarField::<i64>::new(-42i64);
let exp_bytes = (-42i64).to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_f32() {
let field = ScalarField::<f32>::new(42.5f32);
let exp_bytes = 42.5f32.to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_f64() {
let field = ScalarField::<f64>::new(42.5f64);
let exp_bytes = 42.5f64.to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_time_difference() {
let value = TimeDifference::new(42, 1234);
let field = ScalarField::<TimeDifference>::new(value);
let exp_bytes = value.to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_scalar_field_time_of_day() {
let value = TimeOfDay::new(42, 1234);
let field = ScalarField::<TimeOfDay>::new(value);
let exp_bytes = value.to_le_bytes();
sub_read_test_helper(&field, &exp_bytes);
}
#[test]
fn test_byte_field() {
const N: usize = 10;
let field = ByteField::new([0; N]);
let write_data = Vec::from_iter(0u8..N as u8);
field.write(&write_data).unwrap();
sub_read_test_helper(&field, &write_data);
}
#[test]
fn test_null_term_byte_field() {
let field = NullTermByteField::new([0; 10]);
field.write(&[1, 2, 3, 4, 5, 6, 7, 8, 9, 10]).unwrap();
sub_read_test_helper(&field, &[1, 2, 3, 4, 5, 6, 7, 8, 9, 10]);
field.write(&[1, 2, 3, 4]).unwrap();
sub_read_test_helper(&field, &[1, 2, 3, 4]);
}
#[test]
fn test_const_field() {
let field = ConstField::new([1, 2, 3, 4, 5]);
sub_read_test_helper(&field, &[1, 2, 3, 4, 5]);
}
#[test]
fn test_const_byte_ref_field() {
let field = ConstByteRefField::new(&[1, 2, 3, 4, 5]);
sub_read_test_helper(&field, &[1, 2, 3, 4, 5]);
}
}