diff --git a/Cargo.lock b/Cargo.lock index 4cb11cb9..5f86fe03 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2522,7 +2522,7 @@ dependencies = [ [[package]] name = "ros2-client" version = "0.5.2" -source = "git+https://github.com/dora-rs/ros2-client.git?branch=deserialize-seed#2f4f3ce1635745634b9948541acd0c76aff2aa92" +source = "git+https://github.com/dora-rs/ros2-client.git?branch=deserialize-seed#3cab61b9d877322bd9da5779764b8a8e7d3f4030" dependencies = [ "bytes", "cdr-encoding-size", @@ -2570,7 +2570,7 @@ dependencies = [ [[package]] name = "rustdds" version = "0.8.4" -source = "git+https://github.com/dora-rs/RustDDS.git?branch=deserialize-seed#25de36d50d6d1d9986067f885d66556db0f625a2" +source = "git+https://github.com/dora-rs/RustDDS.git?branch=deserialize-seed#f49cafcf6448767a26f0b8788856cb1857a254fc" dependencies = [ "bit-vec", "byteorder", diff --git a/python/src/typed/deserialize.rs b/python/src/typed/deserialize.rs index 0148fd93..74d4b7c4 100644 --- a/python/src/typed/deserialize.rs +++ b/python/src/typed/deserialize.rs @@ -1,15 +1,18 @@ +use super::TypeInfo; use arrow::{ array::{ - make_array, ArrayData, BooleanBuilder, Float32Builder, Float64Builder, Int64Builder, - NullArray, StringBuilder, StructArray, UInt64Builder, + make_array, Array, ArrayData, BooleanBuilder, Float32Builder, Float64Builder, Int16Builder, + Int32Builder, Int64Builder, Int8Builder, ListArray, NullArray, StringBuilder, StructArray, + UInt16Builder, UInt32Builder, UInt64Builder, UInt8Builder, }, - datatypes::{DataType, Fields}, + buffer::OffsetBuffer, + compute::concat, + datatypes::{DataType, Field, Fields}, }; use core::fmt; -use std::ops::Deref; - -use super::TypeInfo; +use std::{ops::Deref, sync::Arc}; +#[derive(Debug, Clone, PartialEq)] pub struct Ros2Value(ArrayData); impl Deref for Ros2Value { @@ -38,7 +41,8 @@ impl<'de> serde::de::DeserializeSeed<'de> for TypedDeserializer { where D: serde::Deserializer<'de>, { - let value = match self.type_info.fields { + let data_type = self.type_info.data_type; + let value = match data_type.clone() { DataType::Struct(fields) => { /// Serde requires that struct and field names are known at /// compile time with a `'static` lifetime, which is not @@ -60,12 +64,29 @@ impl<'de> serde::de::DeserializeSeed<'de> for TypedDeserializer { }, ) } + DataType::List(field) => deserializer.deserialize_seq(ListVisitor { + field, + defaults: self.type_info.defaults, + }), + DataType::UInt8 => deserializer.deserialize_u8(PrimitiveValueVisitor), + DataType::UInt16 => deserializer.deserialize_u16(PrimitiveValueVisitor), + DataType::UInt32 => deserializer.deserialize_u32(PrimitiveValueVisitor), + DataType::UInt64 => deserializer.deserialize_u64(PrimitiveValueVisitor), + DataType::Int8 => deserializer.deserialize_i8(PrimitiveValueVisitor), + DataType::Int16 => deserializer.deserialize_i16(PrimitiveValueVisitor), DataType::Int32 => deserializer.deserialize_i32(PrimitiveValueVisitor), + DataType::Int64 => deserializer.deserialize_i64(PrimitiveValueVisitor), DataType::Float32 => deserializer.deserialize_f32(PrimitiveValueVisitor), DataType::Float64 => deserializer.deserialize_f64(PrimitiveValueVisitor), DataType::Utf8 => deserializer.deserialize_str(PrimitiveValueVisitor), _ => todo!(), }?; + + debug_assert!( + value.data_type() == &data_type, + "Datatype does not correspond to default data type.\n Expected: {:#?} \n but got: {:#?}, with value: {:#?}", data_type, value.data_type(), value + ); + Ok(Ros2Value(value)) } } @@ -89,6 +110,31 @@ impl<'de> serde::de::Visitor<'de> for PrimitiveValueVisitor { Ok(array.finish().into()) } + fn visit_i8(self, u: i8) -> Result + where + E: serde::de::Error, + { + let mut array = Int8Builder::new(); + array.append_value(u); + Ok(array.finish().into()) + } + + fn visit_i16(self, u: i16) -> Result + where + E: serde::de::Error, + { + let mut array = Int16Builder::new(); + array.append_value(u); + Ok(array.finish().into()) + } + fn visit_i32(self, u: i32) -> Result + where + E: serde::de::Error, + { + let mut array = Int32Builder::new(); + array.append_value(u); + Ok(array.finish().into()) + } fn visit_i64(self, i: i64) -> Result where E: serde::de::Error, @@ -98,6 +144,30 @@ impl<'de> serde::de::Visitor<'de> for PrimitiveValueVisitor { Ok(array.finish().into()) } + fn visit_u8(self, u: u8) -> Result + where + E: serde::de::Error, + { + let mut array = UInt8Builder::new(); + array.append_value(u); + Ok(array.finish().into()) + } + fn visit_u16(self, u: u16) -> Result + where + E: serde::de::Error, + { + let mut array = UInt16Builder::new(); + array.append_value(u); + Ok(array.finish().into()) + } + fn visit_u32(self, u: u32) -> Result + where + E: serde::de::Error, + { + let mut array = UInt32Builder::new(); + array.append_value(u); + Ok(array.finish().into()) + } fn visit_u64(self, u: u64) -> Result where E: serde::de::Error, @@ -179,27 +249,162 @@ impl<'de> serde::de::Visitor<'de> for StructVisitor { let mut fields = vec![]; let defaults: StructArray = self.defaults.clone().into(); for field in self.fields.iter() { + let default = match defaults.column_by_name(field.name()) { + Some(value) => value.clone(), + None => { + return Err(serde::de::Error::custom(format!( + "missing field {} for deserialization", + &field.name() + ))) + } + }; let value = match data.next_element_seed(TypedDeserializer { type_info: TypeInfo { - fields: field.data_type().clone(), - defaults: self.defaults.clone(), + data_type: field.data_type().clone(), + defaults: default.to_data(), }, })? { Some(value) => make_array(value.0), - None => match defaults.column_by_name(field.name()) { - Some(value) => value.clone(), - None => { - return Err(serde::de::Error::custom(format!( - "missing field {} for deserialization", - &field.name() - ))) - } - }, + None => default, }; - fields.push((field.clone(), value)); + fields.push(( + // Recreate a new field as List(UInt8) can be converted to UInt8 + Arc::new(Field::new(field.name(), value.data_type().clone(), true)), + value, + )); } let struct_array: StructArray = fields.into(); + Ok(struct_array.into()) } } + +struct ListVisitor { + field: Arc, + defaults: ArrayData, +} + +impl<'de> serde::de::Visitor<'de> for ListVisitor { + type Value = ArrayData; + + fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result { + formatter.write_str("an array encoded as sequence") + } + + fn visit_seq(self, mut data: A) -> Result + where + A: serde::de::SeqAccess<'de>, + { + let data = match self.field.data_type().clone() { + DataType::UInt8 => { + let mut array = UInt8Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::UInt16 => { + let mut array = UInt16Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::UInt32 => { + let mut array = UInt32Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::UInt64 => { + let mut array = UInt64Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::Int8 => { + let mut array = Int8Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::Int16 => { + let mut array = Int16Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::Int32 => { + let mut array = Int32Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::Int64 => { + let mut array = Int64Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::Float32 => { + let mut array = Float32Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::Float64 => { + let mut array = Float64Builder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + DataType::Utf8 => { + let mut array = StringBuilder::new(); + while let Some(value) = data.next_element::()? { + array.append_value(value); + } + Ok(array.finish().into()) + } + _ => { + let mut buffer = vec![]; + while let Some(value) = data.next_element_seed(TypedDeserializer { + type_info: TypeInfo { + data_type: self.field.data_type().clone(), + defaults: self.defaults.clone(), + }, + })? { + let element = make_array(value.0); + buffer.push(element); + } + + concat( + buffer + .iter() + .map(|data| data.as_ref()) + .collect::>() + .as_slice(), + ) + .map(|op| op.to_data()) + } + }; + + if let Ok(values) = data { + let offsets = OffsetBuffer::new(vec![0, values.len() as i32].into()); + + let array = + ListArray::new(self.field, offsets.clone(), make_array(values), None).to_data(); + Ok(array) + } else { + Ok(self.defaults) // TODO: Better handle deserialization error + } + } +} diff --git a/python/src/typed/mod.rs b/python/src/typed/mod.rs index 7a219fe7..4db2103c 100644 --- a/python/src/typed/mod.rs +++ b/python/src/typed/mod.rs @@ -1,9 +1,11 @@ use arrow::{ array::{ - make_array, Array, ArrayData, BooleanArray, Float32Array, Float64Array, Int16Array, - Int32Array, Int64Array, Int8Array, StringArray, StructArray, UInt16Array, UInt32Array, - UInt64Array, UInt8Array, + make_array, Array, ArrayData, ArrayRef, BooleanArray, Float32Array, Float64Array, + Int16Array, Int32Array, Int64Array, Int8Array, StringArray, StructArray, UInt16Array, + UInt32Array, UInt64Array, UInt8Array, }, + buffer::Buffer, + compute::concat, datatypes::{DataType, Field}, }; use dora_ros2_bridge_msg_gen::types::{ @@ -20,7 +22,7 @@ pub mod serialize; #[derive(Debug, Clone, PartialEq)] pub struct TypeInfo { - fields: DataType, + data_type: DataType, defaults: ArrayData, } @@ -38,85 +40,34 @@ pub fn for_message( .members .iter() .map(|m| { + let default = make_array(default_for_member(m, package_name, messages)?); Result::<_, eyre::Report>::Ok(( Arc::new(Field::new( m.name.clone(), - type_info_for_member(m, package_name, messages)?, - false, + default.data_type().clone(), + true, )), - make_array(default_for_member(m, package_name, messages)?), + default, )) }) .collect::>()?; + let default_struct: StructArray = default_struct_vec.into(); + Ok(TypeInfo { - fields: default_struct.data_type().clone(), + data_type: default_struct.data_type().clone(), defaults: default_struct.into(), }) } -fn type_info_for_member( - m: &dora_ros2_bridge_msg_gen::types::Member, - package_name: &str, - messages: &HashMap>, -) -> eyre::Result { - let empty = HashMap::new(); - let package_messages = messages.get(package_name).unwrap_or(&empty); - Ok(match &m.r#type { - MemberType::NestableType(t) => match t { - NestableType::BasicType(t) => match t { - BasicType::I8 => todo!(), - BasicType::I16 => todo!(), - BasicType::I32 => DataType::Int32, - BasicType::I64 => todo!(), - BasicType::U8 => todo!(), - BasicType::U16 => todo!(), - BasicType::U32 => todo!(), - BasicType::U64 => todo!(), - BasicType::F32 => DataType::Float32, - BasicType::F64 => DataType::Float64, - BasicType::Bool => todo!(), - BasicType::Char => todo!(), - BasicType::Byte => todo!(), - }, - NestableType::NamedType(name) => { - let referenced_message = package_messages - .get(&name.0) - .context("unknown referenced message")?; - for_message(messages, package_name, &referenced_message.name)?.fields - } - NestableType::NamespacedType(t) => for_message(messages, &t.package, &t.name)?.fields, - NestableType::GenericString(_) => DataType::Utf8, - }, - MemberType::Array(array) => match array.value_type { - _ => todo!(), - }, - MemberType::Sequence(sequence) => match &sequence.value_type { - NestableType::NamedType(name) => { - let referenced_message = package_messages - .get(&name.0) - .context("unknown referenced message")?; - for_message(messages, package_name, &referenced_message.name)?.fields - } - _ => todo!(), - }, - MemberType::BoundedSequence(_) => { - todo!() - } - }) -} - pub fn default_for_member( m: &dora_ros2_bridge_msg_gen::types::Member, package_name: &str, messages: &HashMap>, ) -> eyre::Result { - let empty = HashMap::new(); - let package_messages = messages.get(package_name).unwrap_or(&empty); - let value = match &m.r#type { MemberType::NestableType(t) => match t { - t @ NestableType::BasicType(_) | t @ NestableType::GenericString(_) => match &m + NestableType::BasicType(_) | NestableType::GenericString(_) => match &m .default .as_deref() { @@ -126,36 +77,40 @@ pub fn default_for_member( Some(_) => eyre::bail!( "there should be only a single default value for non-sequence types" ), - None => default_for_basic_type(t), + None => default_for_nestable_type(t, package_name, messages)?, }, - NestableType::NamedType(name) => { + NestableType::NamedType(_) => { if m.default.is_some() { eyre::bail!("default values for nested types are not supported") } else { - let referenced_message = package_messages - .get(&name.0) - .context("unknown referenced message")?; - - default_for_referenced_message(referenced_message, package_name, messages)? + default_for_nestable_type(t, package_name, messages)? } } - NestableType::NamespacedType(t) => { - let referenced_package_messages = messages.get(&t.package).unwrap_or(&empty); - let referenced_message = referenced_package_messages - .get(&t.name) - .context("unknown referenced message")?; - default_for_referenced_message(referenced_message, &t.package, messages)? + NestableType::NamespacedType(_) => { + default_for_nestable_type(t, package_name, messages)? } }, - MemberType::Array(_) | MemberType::Sequence(_) | MemberType::BoundedSequence(_) => { - todo!() + MemberType::Array(array) => { + list_default_values(m, &array.value_type, package_name, messages)? + } + MemberType::Sequence(seq) => { + list_default_values(m, &seq.value_type, package_name, messages)? + } + MemberType::BoundedSequence(seq) => { + list_default_values(m, &seq.value_type, package_name, messages)? } }; Ok(value) } -fn default_for_basic_type(t: &NestableType) -> ArrayData { - match t { +fn default_for_nestable_type( + t: &NestableType, + package_name: &str, + messages: &HashMap>, +) -> Result { + let empty = HashMap::new(); + let package_messages = messages.get(package_name).unwrap_or(&empty); + let array = match t { NestableType::BasicType(t) => match t { BasicType::I8 => Int8Array::from(vec![0]).into(), BasicType::I16 => Int16Array::from(vec![0]).into(), @@ -168,13 +123,28 @@ fn default_for_basic_type(t: &NestableType) -> ArrayData { BasicType::F32 => Float32Array::from(vec![0.]).into(), BasicType::F64 => Float64Array::from(vec![0.]).into(), BasicType::Char => StringArray::from(vec![""]).into(), - BasicType::Byte => UInt8Array::from(vec![] as Vec).into(), + BasicType::Byte => UInt8Array::from(vec![0u8] as Vec).into(), BasicType::Bool => BooleanArray::from(vec![false]).into(), }, NestableType::GenericString(_) => StringArray::from(vec![""]).into(), - _ => todo!(), - } + NestableType::NamedType(name) => { + let referenced_message = package_messages + .get(&name.0) + .context("unknown referenced message")?; + + default_for_referenced_message(referenced_message, package_name, messages)? + } + NestableType::NamespacedType(t) => { + let referenced_package_messages = messages.get(&t.package).unwrap_or(&empty); + let referenced_message = referenced_package_messages + .get(&t.name) + .context("unknown referenced message")?; + default_for_referenced_message(referenced_message, &t.package, messages)? + } + }; + Ok(array) } + fn preset_default_for_basic_type(t: &NestableType, preset: &str) -> Result { Ok(match t { NestableType::BasicType(t) => match t { @@ -244,7 +214,7 @@ fn default_for_referenced_message( Arc::new(Field::new( m.name.clone(), default.data_type().clone(), - false, + true, )), make_array(default), )) @@ -254,3 +224,60 @@ fn default_for_referenced_message( let struct_array: StructArray = fields.into(); Ok(struct_array.into()) } + +fn list_default_values( + m: &dora_ros2_bridge_msg_gen::types::Member, + value_type: &NestableType, + package_name: &str, + messages: &HashMap>, +) -> Result { + let defaults = match &m.default.as_deref() { + Some([]) => eyre::bail!("empty default value not supported"), + Some(defaults) => { + let raw_array: Vec> = defaults + .iter() + .map(|default| { + preset_default_for_basic_type(value_type, &default) + .with_context(|| format!("failed to parse default value for `{}`", m.name)) + .map(|data| make_array(data)) + }) + .collect::>()?; + let default_values = concat( + raw_array + .iter() + .map(|data| data.as_ref()) + .collect::>() + .as_slice(), + ) + .context("Failed to concatenate default list value")?; + default_values.to_data() + } + None => { + let default_nested_type = + default_for_nestable_type(&value_type, package_name, messages)?; + if false { + //let NestableType::BasicType(_t) = seq.value_type { + default_nested_type.into() + } else { + let value_offsets = Buffer::from_slice_ref([0i64, 1]); + + let list_data_type = DataType::List(Arc::new(Field::new( + &m.name, + default_nested_type.data_type().clone(), + true, + ))); + // Construct a list array from the above two + let array = ArrayData::builder(list_data_type) + .len(1) + .add_buffer(value_offsets.clone()) + .add_child_data(default_nested_type.clone()) + .build() + .unwrap(); + + array.into() + } + } + }; + + Ok(defaults) +} diff --git a/python/src/typed/serialize.rs b/python/src/typed/serialize.rs index 36329cde..1d02f64d 100644 --- a/python/src/typed/serialize.rs +++ b/python/src/typed/serialize.rs @@ -1,10 +1,19 @@ use arrow::array::ArrayData; use arrow::array::Float32Array; use arrow::array::Float64Array; +use arrow::array::Int16Array; use arrow::array::Int32Array; +use arrow::array::Int64Array; +use arrow::array::Int8Array; +use arrow::array::ListArray; use arrow::array::StringArray; use arrow::array::StructArray; +use arrow::array::UInt16Array; +use arrow::array::UInt32Array; +use arrow::array::UInt64Array; +use arrow::array::UInt8Array; use arrow::datatypes::DataType; +use serde::ser::SerializeSeq; use serde::ser::SerializeStruct; use super::TypeInfo; @@ -21,11 +30,46 @@ impl serde::Serialize for TypedValue<'_> { S: serde::Serializer, { match &self.value.data_type() { + DataType::UInt8 => { + let uint_array: UInt8Array = self.value.clone().into(); + let number = uint_array.value(0); + serializer.serialize_u8(number) + } + DataType::UInt16 => { + let uint_array: UInt16Array = self.value.clone().into(); + let number = uint_array.value(0); + serializer.serialize_u16(number) + } + DataType::UInt32 => { + let uint_array: UInt32Array = self.value.clone().into(); + let number = uint_array.value(0); + serializer.serialize_u32(number) + } + DataType::UInt64 => { + let uint_array: UInt64Array = self.value.clone().into(); + let number = uint_array.value(0); + serializer.serialize_u64(number) + } + DataType::Int8 => { + let int_array: Int8Array = self.value.clone().into(); + let number = int_array.value(0); + serializer.serialize_i8(number) + } + DataType::Int16 => { + let int_array: Int16Array = self.value.clone().into(); + let number = int_array.value(0); + serializer.serialize_i16(number) + } DataType::Int32 => { let int_array: Int32Array = self.value.clone().into(); let number = int_array.value(0); serializer.serialize_i32(number) } + DataType::Int64 => { + let int_array: Int64Array = self.value.clone().into(); + let number = int_array.value(0); + serializer.serialize_i64(number) + } DataType::Float32 => { let int_array: Float32Array = self.value.clone().into(); let number = int_array.value(0); @@ -41,7 +85,34 @@ impl serde::Serialize for TypedValue<'_> { let string = int_array.value(0); serializer.serialize_str(string) } - DataType::List(_field_ref) => todo!(), + DataType::List(_field) => { + let list_array: ListArray = self.value.clone().into(); + if let DataType::List(field) = self.type_info.data_type.clone() { + let values = list_array.values(); + let mut s = serializer.serialize_seq(Some(values.len()))?; + for value in list_array.iter() { + let value = match value { + Some(value) => value.to_data(), + None => { + return Err(serde::ser::Error::custom(format!( + "Value in ListArray is null and not yet supported", + ))) + } + }; + + s.serialize_element(&TypedValue { + value: &value, + type_info: &TypeInfo { + data_type: field.data_type().clone(), + defaults: self.type_info.defaults.clone(), + }, + })?; + } + s.end() + } else { + return Err(serde::ser::Error::custom(format!("Wrong fields type",))); + } + } DataType::Struct(_fields) => { /// Serde requires that struct and field names are known at /// compile time with a `'static` lifetime, which is not @@ -55,7 +126,7 @@ impl serde::Serialize for TypedValue<'_> { const DUMMY_FIELD_NAME: &str = "field"; let struct_array: StructArray = self.value.clone().into(); - if let DataType::Struct(fields) = self.type_info.fields.clone() { + if let DataType::Struct(fields) = self.type_info.data_type.clone() { let mut s = serializer.serialize_struct(DUMMY_STRUCT_NAME, fields.len())?; let defaults: StructArray = self.type_info.defaults.clone().into(); for field in fields.iter() { @@ -78,7 +149,7 @@ impl serde::Serialize for TypedValue<'_> { &TypedValue { value: &field_value, type_info: &TypeInfo { - fields: field.data_type().clone(), + data_type: field.data_type().clone(), defaults: default, }, },