From 811186f99db91ce9adf82b92eadfc6c3ea5bc923 Mon Sep 17 00:00:00 2001 From: haixuanTao Date: Wed, 23 Aug 2023 12:17:24 +0200 Subject: [PATCH 1/7] Remove `todo` by adding basic arrow type --- python/src/typed/deserialize.rs | 66 +++++++++++++++++++++++++++++++-- python/src/typed/serialize.rs | 43 +++++++++++++++++++++ 2 files changed, 105 insertions(+), 4 deletions(-) diff --git a/python/src/typed/deserialize.rs b/python/src/typed/deserialize.rs index 0148fd93..ad276f43 100644 --- a/python/src/typed/deserialize.rs +++ b/python/src/typed/deserialize.rs @@ -1,12 +1,14 @@ 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, NullArray, StringBuilder, StructArray, + UInt16Builder, UInt32Builder, UInt64Builder, UInt8Builder, }, - datatypes::{DataType, Fields}, + compute::concat, + datatypes::{DataType, Field, Fields}, }; use core::fmt; -use std::ops::Deref; +use std::{ops::Deref, sync::Arc}; use super::TypeInfo; @@ -60,7 +62,14 @@ impl<'de> serde::de::DeserializeSeed<'de> for TypedDeserializer { }, ) } + 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), @@ -89,6 +98,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 +132,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, diff --git a/python/src/typed/serialize.rs b/python/src/typed/serialize.rs index 36329cde..7571ed16 100644 --- a/python/src/typed/serialize.rs +++ b/python/src/typed/serialize.rs @@ -1,9 +1,17 @@ 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::SerializeStruct; @@ -21,11 +29,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); From 44788759353b997de79c8797fc3c99abf9e53062 Mon Sep 17 00:00:00 2001 From: haixuanTao Date: Wed, 23 Aug 2023 12:21:47 +0200 Subject: [PATCH 2/7] Add `List` deserializer --- python/src/typed/deserialize.rs | 96 ++++++++++++++++++++++++++++----- 1 file changed, 82 insertions(+), 14 deletions(-) diff --git a/python/src/typed/deserialize.rs b/python/src/typed/deserialize.rs index ad276f43..f5f9a6e9 100644 --- a/python/src/typed/deserialize.rs +++ b/python/src/typed/deserialize.rs @@ -1,9 +1,10 @@ use arrow::{ array::{ make_array, Array, ArrayData, BooleanBuilder, Float32Builder, Float64Builder, Int16Builder, - Int32Builder, Int64Builder, Int8Builder, NullArray, StringBuilder, StructArray, + Int32Builder, Int64Builder, Int8Builder, ListArray, NullArray, StringBuilder, StructArray, UInt16Builder, UInt32Builder, UInt64Builder, UInt8Builder, }, + buffer::OffsetBuffer, compute::concat, datatypes::{DataType, Field, Fields}, }; @@ -12,6 +13,7 @@ use std::{ops::Deref, sync::Arc}; use super::TypeInfo; +#[derive(Debug, Clone, PartialEq)] pub struct Ros2Value(ArrayData); impl Deref for Ros2Value { @@ -40,7 +42,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 @@ -62,6 +65,10 @@ 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), @@ -237,27 +244,88 @@ 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 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); + } + + if let Ok(array) = concat( + buffer + .iter() + .map(|data| data.as_ref()) + .collect::>() + .as_slice(), + ) { + let values = array; //.to_data(); + + let offsets = OffsetBuffer::new(vec![0, values.len() as i32].into()); + + let field = Arc::new(Field::new( + self.field.name(), + values.data_type().clone(), + false, + )); + let array = ListArray::new(field.clone(), offsets.clone(), values.clone(), None); + Ok(array.to_data()) + } else { + Ok(self.defaults) // TODO: Better handle deserialization error + //return Err(serde::de::Error::custom(format!( + //"Could not parse ROS2 list of values", + //))); + } + } +} From d4207595fd46553f04b612a70236d1cf721d8576 Mon Sep 17 00:00:00 2001 From: haixuanTao Date: Wed, 23 Aug 2023 12:23:02 +0200 Subject: [PATCH 3/7] Refactor default and type info to enable recursive typing --- python/src/typed/mod.rs | 278 ++++++++++++++++++++++++++++++---------- 1 file changed, 208 insertions(+), 70 deletions(-) diff --git a/python/src/typed/mod.rs b/python/src/typed/mod.rs index 7a219fe7..eee31a60 100644 --- a/python/src/typed/mod.rs +++ b/python/src/typed/mod.rs @@ -4,6 +4,7 @@ use arrow::{ Int32Array, Int64Array, Int8Array, StringArray, StructArray, UInt16Array, UInt32Array, UInt64Array, UInt8Array, }, + buffer::Buffer, datatypes::{DataType, Field}, }; use dora_ros2_bridge_msg_gen::types::{ @@ -20,7 +21,7 @@ pub mod serialize; #[derive(Debug, Clone, PartialEq)] pub struct TypeInfo { - fields: DataType, + data_type: DataType, defaults: ArrayData, } @@ -38,19 +39,22 @@ 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(), }) } @@ -60,50 +64,78 @@ fn type_info_for_member( package_name: &str, messages: &HashMap>, ) -> eyre::Result { + Ok(match &m.r#type { + MemberType::NestableType(t) => type_info_for_nestable_type(t, package_name, messages)?, + MemberType::Array(array) => Field::new_list( + &m.name, + Field::new( + &m.name, + type_info_for_nestable_type(&array.value_type, package_name, messages)?, + true, + ), + true, + ) + .data_type() + .clone(), + MemberType::Sequence(sequence) => Field::new_list( + &m.name, + Field::new( + &m.name, + type_info_for_nestable_type(&sequence.value_type, package_name, messages)?, + true, + ), + true, + ) + .data_type() + .clone(), + MemberType::BoundedSequence(sequence) => Field::new_list( + &m.name, + Field::new( + &m.name, + type_info_for_nestable_type(&sequence.value_type, package_name, messages)?, + true, + ), + true, + ) + .data_type() + .clone(), + }) +} + +fn type_info_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); - 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!(), + let data_type = match t { + NestableType::BasicType(t) => match t { + BasicType::I8 => DataType::Int8, + BasicType::I16 => DataType::Int16, + BasicType::I32 => DataType::Int32, + BasicType::I64 => DataType::Int64, + BasicType::U8 => DataType::UInt8, + BasicType::U16 => DataType::UInt16, + BasicType::U32 => DataType::UInt32, + BasicType::U64 => DataType::UInt64, + BasicType::F32 => DataType::Float32, + BasicType::F64 => DataType::Float64, + BasicType::Bool => DataType::Boolean, + BasicType::Char => DataType::Utf8, + BasicType::Byte => DataType::UInt8, }, - MemberType::BoundedSequence(_) => { - todo!() + NestableType::NamedType(name) => { + let referenced_message = package_messages + .get(&name.0) + .context("unknown referenced message")?; + for_message(messages, package_name, &referenced_message.name)?.data_type } - }) + NestableType::NamespacedType(t) => for_message(messages, &t.package, &t.name)?.data_type, + NestableType::GenericString(_) => DataType::Utf8, + }; + + Ok(data_type) } pub fn default_for_member( @@ -111,12 +143,9 @@ pub fn default_for_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 +155,130 @@ 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_nestable_type(t, package_name, messages)? + } + } + NestableType::NamespacedType(_) => { + default_for_nestable_type(t, package_name, messages)? + } + }, + MemberType::Array(array) => match &m.default.as_deref() { + Some([]) => eyre::bail!("empty default value not supported"), + Some([default]) => preset_default_for_basic_type(&array.value_type, &default) + .with_context(|| format!("failed to parse default value for `{}`", m.name))?, + Some(_) => { + eyre::bail!("there should be only a single default value for non-sequence types") + } + None => { + let default_nested_type = + default_for_nestable_type(&array.value_type, package_name, messages)?; + if false { + //let NestableType::BasicType(_t) = array.value_type { + default_nested_type.into() + } else { + let value_offsets = Buffer::from_slice_ref([0i64]); - default_for_referenced_message(referenced_message, package_name, messages)? + 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() } } - 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)? + }, + MemberType::Sequence(seq) => match &m.default.as_deref() { + Some([]) => eyre::bail!("empty default value not supported"), + Some([default]) => preset_default_for_basic_type(&seq.value_type, &default) + .with_context(|| format!("failed to parse default value for `{}`", m.name))?, + Some(_) => { + eyre::bail!("there should be only a single default value for non-sequence types") + } + None => { + let default_nested_type = + default_for_nestable_type(&seq.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]); + + 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() + } + } + }, + MemberType::BoundedSequence(seq) => match &m.default.as_deref() { + Some([]) => eyre::bail!("empty default value not supported"), + Some([default]) => preset_default_for_basic_type(&seq.value_type, &default) + .with_context(|| format!("failed to parse default value for `{}`", m.name))?, + Some(_) => { + eyre::bail!("there should be only a single default value for non-sequence types") + } + None => { + let default_nested_type = + default_for_nestable_type(&seq.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]); + + 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() + } } }, - MemberType::Array(_) | MemberType::Sequence(_) | MemberType::BoundedSequence(_) => { - todo!() - } }; 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 +291,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 +382,7 @@ fn default_for_referenced_message( Arc::new(Field::new( m.name.clone(), default.data_type().clone(), - false, + true, )), make_array(default), )) From 2fd654393bbc38a385fe0d7d12148d34b14a1315 Mon Sep 17 00:00:00 2001 From: haixuanTao Date: Wed, 23 Aug 2023 12:24:13 +0200 Subject: [PATCH 4/7] Serialize `List` for arrow --- Cargo.lock | 379 ++++++++++++++++++++++++++-------- python/src/typed/serialize.rs | 34 ++- 2 files changed, 324 insertions(+), 89 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 13de3e4c..101794dd 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -122,19 +122,41 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "73bdeeaf5bbeeb40c6e14849520b379cd22d8f605433e7c12a1d550adf8c4a06" dependencies = [ "ahash", - "arrow-arith", - "arrow-array", - "arrow-buffer", - "arrow-cast", - "arrow-csv", - "arrow-data", - "arrow-ipc", - "arrow-json", - "arrow-ord", - "arrow-row", - "arrow-schema", - "arrow-select", - "arrow-string", + "arrow-arith 35.0.0", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-cast 35.0.0", + "arrow-csv 35.0.0", + "arrow-data 35.0.0", + "arrow-ipc 35.0.0", + "arrow-json 35.0.0", + "arrow-ord 35.0.0", + "arrow-row 35.0.0", + "arrow-schema 35.0.0", + "arrow-select 35.0.0", + "arrow-string 35.0.0", +] + +[[package]] +name = "arrow" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7104b9e9761613ae92fe770c741d6bbf1dbc791a0fe204400aebdd429875741" +dependencies = [ + "ahash", + "arrow-arith 45.0.0", + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-cast 45.0.0", + "arrow-csv 45.0.0", + "arrow-data 45.0.0", + "arrow-ipc 45.0.0", + "arrow-json 45.0.0", + "arrow-ord 45.0.0", + "arrow-row 45.0.0", + "arrow-schema 45.0.0", + "arrow-select 45.0.0", + "arrow-string 45.0.0", "pyo3", ] @@ -144,10 +166,25 @@ version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "24a945eab89f800ab870b848a6105b464638271565d4ac80c439a349f6dae349" dependencies = [ - "arrow-array", - "arrow-buffer", - "arrow-data", - "arrow-schema", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", + "chrono", + "half", + "num", +] + +[[package]] +name = "arrow-arith" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38e597a8e8efb8ff52c50eaf8f4d85124ce3c1bf20fab82f476d73739d9ab1c2" +dependencies = [ + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", "chrono", "half", "num", @@ -160,15 +197,31 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "43489bbff475545b78b0e20bde1d22abd6c99e54499839f9e815a2fa5134a51b" dependencies = [ "ahash", - "arrow-buffer", - "arrow-data", - "arrow-schema", + "arrow-buffer 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", "chrono", "half", "hashbrown 0.13.2", "num", ] +[[package]] +name = "arrow-array" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a86d9c1473db72896bd2345ebb6b8ad75b8553ba390875c76708e8dc5c5492d" +dependencies = [ + "ahash", + "arrow-buffer 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", + "chrono", + "half", + "hashbrown 0.14.0", + "num", +] + [[package]] name = "arrow-buffer" version = "35.0.0" @@ -179,33 +232,79 @@ dependencies = [ "num", ] +[[package]] +name = "arrow-buffer" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "234b3b1c8ed00c874bf95972030ac4def6f58e02ea5a7884314388307fb3669b" +dependencies = [ + "half", + "num", +] + [[package]] name = "arrow-cast" version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b30c01f06172d3e8306fcc885ee97dff55ba8d48dbc2ef5fee2e75b55b8170c4" dependencies = [ - "arrow-array", - "arrow-buffer", - "arrow-data", - "arrow-schema", - "arrow-select", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", + "arrow-select 35.0.0", "chrono", "lexical-core", "num", ] +[[package]] +name = "arrow-cast" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22f61168b853c7faea8cea23a2169fdff9c82fb10ae5e2c07ad1cab8f6884931" +dependencies = [ + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", + "arrow-select 45.0.0", + "chrono", + "half", + "lexical-core", + "num", +] + [[package]] name = "arrow-csv" version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "86059a1270d5fe268283447af53bda33f4226fbe9e3f3d109a26011dc97fdcf7" dependencies = [ - "arrow-array", - "arrow-buffer", - "arrow-cast", - "arrow-data", - "arrow-schema", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-cast 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", + "chrono", + "csv", + "csv-core", + "lazy_static", + "lexical-core", + "regex", +] + +[[package]] +name = "arrow-csv" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10b545c114d9bf8569c84d2fbe2020ac4eea8db462c0a37d0b65f41a90d066fe" +dependencies = [ + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-cast 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", "chrono", "csv", "csv-core", @@ -220,8 +319,20 @@ version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "19c7787c6cdbf9539b1ffb860bfc18c5848926ec3d62cbd52dc3b1ea35c874fd" dependencies = [ - "arrow-buffer", - "arrow-schema", + "arrow-buffer 35.0.0", + "arrow-schema 35.0.0", + "half", + "num", +] + +[[package]] +name = "arrow-data" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6b6852635e7c43e5b242841c7470606ff0ee70eef323004cacc3ecedd33dd8f" +dependencies = [ + "arrow-buffer 45.0.0", + "arrow-schema 45.0.0", "half", "num", ] @@ -232,11 +343,25 @@ version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "690167cd0ad8c4444c7bbb573066b94982acb52da554db7b7b2837c0e2dce036" dependencies = [ - "arrow-array", - "arrow-buffer", - "arrow-cast", - "arrow-data", - "arrow-schema", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-cast 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", + "flatbuffers", +] + +[[package]] +name = "arrow-ipc" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a66da9e16aecd9250af0ae9717ae8dd7ea0d8ca5a3e788fe3de9f4ee508da751" +dependencies = [ + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-cast 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", "flatbuffers", ] @@ -246,11 +371,11 @@ version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "28c47c78411bc7b77ab0b931b9f8acc8b1c6d7d9977e5895dbaa4632651e5824" dependencies = [ - "arrow-array", - "arrow-buffer", - "arrow-cast", - "arrow-data", - "arrow-schema", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-cast 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", "chrono", "half", "indexmap 1.9.3", @@ -259,17 +384,52 @@ dependencies = [ "serde_json", ] +[[package]] +name = "arrow-json" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "60ee0f9d8997f4be44a60ee5807443e396e025c23cf14d2b74ce56135cb04474" +dependencies = [ + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-cast 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", + "chrono", + "half", + "indexmap 2.0.0", + "lexical-core", + "num", + "serde", + "serde_json", +] + [[package]] name = "arrow-ord" version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d3945e160dc92c11c4108a13251940ffce418c221eed260c53d9bec4075aa52c" dependencies = [ - "arrow-array", - "arrow-buffer", - "arrow-data", - "arrow-schema", - "arrow-select", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", + "arrow-select 35.0.0", + "num", +] + +[[package]] +name = "arrow-ord" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7fcab05410e6b241442abdab6e1035177dc082bdb6f17049a4db49faed986d63" +dependencies = [ + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", + "arrow-select 45.0.0", + "half", "num", ] @@ -280,21 +440,42 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1951a64d60c37931ee85198e1728ad07e59cda50b9280367911c7d299dfe4bf7" dependencies = [ "ahash", - "arrow-array", - "arrow-buffer", - "arrow-data", - "arrow-schema", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", "half", "hashbrown 0.13.2", ] +[[package]] +name = "arrow-row" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91a847dd9eb0bacd7836ac63b3475c68b2210c2c96d0ec1b808237b973bd5d73" +dependencies = [ + "ahash", + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", + "half", + "hashbrown 0.14.0", +] + [[package]] name = "arrow-schema" version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bf6b26f6a6f8410e3b9531cbd1886399b99842701da77d4b4cf2013f7708f20f" + +[[package]] +name = "arrow-schema" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "54df8c47918eb634c20e29286e69494fdc20cafa5173eb6dad49c7f6acece733" dependencies = [ - "bitflags 1.3.2", + "bitflags 2.3.3", ] [[package]] @@ -303,10 +484,23 @@ version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "83deb30a09afdf654d346092ef03e965f9c05d9b0753164d531f41d290d553f4" dependencies = [ - "arrow-array", - "arrow-buffer", - "arrow-data", - "arrow-schema", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", + "num", +] + +[[package]] +name = "arrow-select" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "941dbe481da043c4bd40c805a19ec2fc008846080c4953171b62bcad5ee5f7fb" +dependencies = [ + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", "num", ] @@ -316,15 +510,31 @@ version = "35.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b7529ba37a8bdc86cb69e00236064d122b3ac80e221270c7abc147fea63a9cee" dependencies = [ - "arrow-array", - "arrow-buffer", - "arrow-data", - "arrow-schema", - "arrow-select", + "arrow-array 35.0.0", + "arrow-buffer 35.0.0", + "arrow-data 35.0.0", + "arrow-schema 35.0.0", + "arrow-select 35.0.0", "regex", "regex-syntax 0.6.29", ] +[[package]] +name = "arrow-string" +version = "45.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "359b2cd9e071d5a3bcf44679f9d85830afebc5b9c98a08019a570a65ae933e0f" +dependencies = [ + "arrow-array 45.0.0", + "arrow-buffer 45.0.0", + "arrow-data 45.0.0", + "arrow-schema 45.0.0", + "arrow-select 45.0.0", + "num", + "regex", + "regex-syntax 0.7.4", +] + [[package]] name = "async-trait" version = "0.1.72" @@ -793,7 +1003,7 @@ version = "0.2.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "aad1a982e6c52b371ca3523ec001def2b5c5b15ca2d0de2c0d23838154baa0ef" dependencies = [ - "arrow", + "arrow 35.0.0", "bincode", "capnp", "dora-core", @@ -873,15 +1083,13 @@ dependencies = [ name = "dora-ros2-bridge-python" version = "0.1.0" dependencies = [ - "arrow", + "arrow 45.0.0", "dora-ros2-bridge", "dora-ros2-bridge-msg-gen", "eyre", "flume", "pyo3", - "pythonize", "serde", - "serde_yaml 0.8.26", ] [[package]] @@ -1726,6 +1934,15 @@ dependencies = [ "autocfg", ] +[[package]] +name = "memoffset" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a634b1c61a95585bd15607c6ab0c4e5b226e695ff2800ba0cdccddf208c406c" +dependencies = [ + "autocfg", +] + [[package]] name = "mime" version = "0.3.17" @@ -2224,15 +2441,15 @@ dependencies = [ [[package]] name = "pyo3" -version = "0.18.3" +version = "0.19.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3b1ac5b3731ba34fdaa9785f8d74d17448cd18f30cf19e0c7e7b1fdb5272109" +checksum = "e681a6cfdc4adcc93b4d3cf993749a4552018ee0a9b65fc0ccfad74352c72a38" dependencies = [ "cfg-if 1.0.0", "eyre", "indoc", "libc", - "memoffset 0.8.0", + "memoffset 0.9.0", "parking_lot", "pyo3-build-config", "pyo3-ffi", @@ -2243,9 +2460,9 @@ dependencies = [ [[package]] name = "pyo3-build-config" -version = "0.18.3" +version = "0.19.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9cb946f5ac61bb61a5014924910d936ebd2b23b705f7a4a3c40b05c720b079a3" +checksum = "076c73d0bc438f7a4ef6fdd0c3bb4732149136abd952b110ac93e4edb13a6ba5" dependencies = [ "once_cell", "target-lexicon", @@ -2253,9 +2470,9 @@ dependencies = [ [[package]] name = "pyo3-ffi" -version = "0.18.3" +version = "0.19.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fd4d7c5337821916ea2a1d21d1092e8443cf34879e53a0ac653fbb98f44ff65c" +checksum = "e53cee42e77ebe256066ba8aa77eff722b3bb91f3419177cf4cd0f304d3284d9" dependencies = [ "libc", "pyo3-build-config", @@ -2263,9 +2480,9 @@ dependencies = [ [[package]] name = "pyo3-macros" -version = "0.18.3" +version = "0.19.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a9d39c55dab3fc5a4b25bbd1ac10a2da452c4aca13bb450f22818a002e29648d" +checksum = "dfeb4c99597e136528c6dd7d5e3de5434d1ceaf487436a3f03b2d56b6fc9efd1" dependencies = [ "proc-macro2", "pyo3-macros-backend", @@ -2275,25 +2492,15 @@ dependencies = [ [[package]] name = "pyo3-macros-backend" -version = "0.18.3" +version = "0.19.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "97daff08a4c48320587b5224cc98d609e3c27b6d437315bd40b605c98eeb5918" +checksum = "947dc12175c254889edc0c02e399476c2f652b4b9ebd123aa655c224de259536" dependencies = [ "proc-macro2", "quote", "syn 1.0.109", ] -[[package]] -name = "pythonize" -version = "0.18.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9a0e1bbcd2a3856284bf4f4ef09ccb1157e9467847792754556f153ea3fe6b42" -dependencies = [ - "pyo3", - "serde", -] - [[package]] name = "quote" version = "1.0.32" diff --git a/python/src/typed/serialize.rs b/python/src/typed/serialize.rs index 7571ed16..1d02f64d 100644 --- a/python/src/typed/serialize.rs +++ b/python/src/typed/serialize.rs @@ -13,6 +13,7 @@ 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; @@ -84,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 @@ -98,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() { @@ -121,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, }, }, From f12f0e339fd2d0781c3dabee5b8b6405ac4a84eb Mon Sep 17 00:00:00 2001 From: haixuanTao Date: Thu, 24 Aug 2023 16:23:51 +0200 Subject: [PATCH 5/7] Remove type info --- python/src/typed/mod.rs | 81 +---------------------------------------- 1 file changed, 1 insertion(+), 80 deletions(-) diff --git a/python/src/typed/mod.rs b/python/src/typed/mod.rs index eee31a60..40b0700f 100644 --- a/python/src/typed/mod.rs +++ b/python/src/typed/mod.rs @@ -59,85 +59,6 @@ pub fn for_message( }) } -fn type_info_for_member( - m: &dora_ros2_bridge_msg_gen::types::Member, - package_name: &str, - messages: &HashMap>, -) -> eyre::Result { - Ok(match &m.r#type { - MemberType::NestableType(t) => type_info_for_nestable_type(t, package_name, messages)?, - MemberType::Array(array) => Field::new_list( - &m.name, - Field::new( - &m.name, - type_info_for_nestable_type(&array.value_type, package_name, messages)?, - true, - ), - true, - ) - .data_type() - .clone(), - MemberType::Sequence(sequence) => Field::new_list( - &m.name, - Field::new( - &m.name, - type_info_for_nestable_type(&sequence.value_type, package_name, messages)?, - true, - ), - true, - ) - .data_type() - .clone(), - MemberType::BoundedSequence(sequence) => Field::new_list( - &m.name, - Field::new( - &m.name, - type_info_for_nestable_type(&sequence.value_type, package_name, messages)?, - true, - ), - true, - ) - .data_type() - .clone(), - }) -} - -fn type_info_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 data_type = match t { - NestableType::BasicType(t) => match t { - BasicType::I8 => DataType::Int8, - BasicType::I16 => DataType::Int16, - BasicType::I32 => DataType::Int32, - BasicType::I64 => DataType::Int64, - BasicType::U8 => DataType::UInt8, - BasicType::U16 => DataType::UInt16, - BasicType::U32 => DataType::UInt32, - BasicType::U64 => DataType::UInt64, - BasicType::F32 => DataType::Float32, - BasicType::F64 => DataType::Float64, - BasicType::Bool => DataType::Boolean, - BasicType::Char => DataType::Utf8, - BasicType::Byte => DataType::UInt8, - }, - NestableType::NamedType(name) => { - let referenced_message = package_messages - .get(&name.0) - .context("unknown referenced message")?; - for_message(messages, package_name, &referenced_message.name)?.data_type - } - NestableType::NamespacedType(t) => for_message(messages, &t.package, &t.name)?.data_type, - NestableType::GenericString(_) => DataType::Utf8, - }; - - Ok(data_type) -} - pub fn default_for_member( m: &dora_ros2_bridge_msg_gen::types::Member, package_name: &str, @@ -248,7 +169,7 @@ pub fn default_for_member( //let NestableType::BasicType(_t) = seq.value_type { default_nested_type.into() } else { - let value_offsets = Buffer::from_slice_ref([0i64]); + let value_offsets = Buffer::from_slice_ref([0i64, 1]); let list_data_type = DataType::List(Arc::new(Field::new( &m.name, From e7e016a8253920123324f44ab4070785b1755789 Mon Sep 17 00:00:00 2001 From: haixuanTao Date: Thu, 24 Aug 2023 16:27:29 +0200 Subject: [PATCH 6/7] Optimize deserializing of basic type --- python/src/typed/deserialize.rs | 141 +++++++++++++++++++++++++------- 1 file changed, 110 insertions(+), 31 deletions(-) diff --git a/python/src/typed/deserialize.rs b/python/src/typed/deserialize.rs index f5f9a6e9..74d4b7c4 100644 --- a/python/src/typed/deserialize.rs +++ b/python/src/typed/deserialize.rs @@ -1,3 +1,4 @@ +use super::TypeInfo; use arrow::{ array::{ make_array, Array, ArrayData, BooleanBuilder, Float32Builder, Float64Builder, Int16Builder, @@ -11,8 +12,6 @@ use arrow::{ use core::fmt; use std::{ops::Deref, sync::Arc}; -use super::TypeInfo; - #[derive(Debug, Clone, PartialEq)] pub struct Ros2Value(ArrayData); @@ -82,6 +81,12 @@ impl<'de> serde::de::DeserializeSeed<'de> for TypedDeserializer { 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)) } } @@ -291,41 +296,115 @@ impl<'de> serde::de::Visitor<'de> for ListVisitor { where A: serde::de::SeqAccess<'de>, { - 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); - } + 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); + } - if let Ok(array) = concat( - buffer - .iter() - .map(|data| data.as_ref()) - .collect::>() - .as_slice(), - ) { - let values = array; //.to_data(); + 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 field = Arc::new(Field::new( - self.field.name(), - values.data_type().clone(), - false, - )); - let array = ListArray::new(field.clone(), offsets.clone(), values.clone(), None); - Ok(array.to_data()) + 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 - //return Err(serde::de::Error::custom(format!( - //"Could not parse ROS2 list of values", - //))); } } } From e55be5c79d431bb4d9fb2ca46a8f7480677a6bad Mon Sep 17 00:00:00 2001 From: haixuanTao Date: Thu, 24 Aug 2023 17:40:43 +0200 Subject: [PATCH 7/7] Compute list default value when multiple default are provided --- python/src/typed/mod.rs | 172 ++++++++++++++++------------------------ 1 file changed, 70 insertions(+), 102 deletions(-) diff --git a/python/src/typed/mod.rs b/python/src/typed/mod.rs index 40b0700f..4db2103c 100644 --- a/python/src/typed/mod.rs +++ b/python/src/typed/mod.rs @@ -1,10 +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::{ @@ -89,105 +90,15 @@ pub fn default_for_member( default_for_nestable_type(t, package_name, messages)? } }, - MemberType::Array(array) => match &m.default.as_deref() { - Some([]) => eyre::bail!("empty default value not supported"), - Some([default]) => preset_default_for_basic_type(&array.value_type, &default) - .with_context(|| format!("failed to parse default value for `{}`", m.name))?, - Some(_) => { - eyre::bail!("there should be only a single default value for non-sequence types") - } - None => { - let default_nested_type = - default_for_nestable_type(&array.value_type, package_name, messages)?; - if false { - //let NestableType::BasicType(_t) = array.value_type { - default_nested_type.into() - } else { - let value_offsets = Buffer::from_slice_ref([0i64]); - - 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() - } - } - }, - MemberType::Sequence(seq) => match &m.default.as_deref() { - Some([]) => eyre::bail!("empty default value not supported"), - Some([default]) => preset_default_for_basic_type(&seq.value_type, &default) - .with_context(|| format!("failed to parse default value for `{}`", m.name))?, - Some(_) => { - eyre::bail!("there should be only a single default value for non-sequence types") - } - None => { - let default_nested_type = - default_for_nestable_type(&seq.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]); - - 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() - } - } - }, - MemberType::BoundedSequence(seq) => match &m.default.as_deref() { - Some([]) => eyre::bail!("empty default value not supported"), - Some([default]) => preset_default_for_basic_type(&seq.value_type, &default) - .with_context(|| format!("failed to parse default value for `{}`", m.name))?, - Some(_) => { - eyre::bail!("there should be only a single default value for non-sequence types") - } - None => { - let default_nested_type = - default_for_nestable_type(&seq.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() - } - } - }, + 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) } @@ -313,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) +}