Browse Source

Merge pull request #9 from dora-rs/arrow-list

Arrow list
tags/v0.2.5-alpha.2
Philipp Oppermann GitHub 2 years ago
parent
commit
c08bb14406
No known key found for this signature in database GPG Key ID: 4AEE18F83AFDEB23
4 changed files with 411 additions and 108 deletions
  1. +2
    -2
      Cargo.lock
  2. +224
    -19
      python/src/typed/deserialize.rs
  3. +111
    -84
      python/src/typed/mod.rs
  4. +74
    -3
      python/src/typed/serialize.rs

+ 2
- 2
Cargo.lock View File

@@ -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",


+ 224
- 19
python/src/typed/deserialize.rs View File

@@ -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<E>(self, u: i8) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
let mut array = Int8Builder::new();
array.append_value(u);
Ok(array.finish().into())
}

fn visit_i16<E>(self, u: i16) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
let mut array = Int16Builder::new();
array.append_value(u);
Ok(array.finish().into())
}
fn visit_i32<E>(self, u: i32) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
let mut array = Int32Builder::new();
array.append_value(u);
Ok(array.finish().into())
}
fn visit_i64<E>(self, i: i64) -> Result<Self::Value, E>
where
E: serde::de::Error,
@@ -98,6 +144,30 @@ impl<'de> serde::de::Visitor<'de> for PrimitiveValueVisitor {
Ok(array.finish().into())
}

fn visit_u8<E>(self, u: u8) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
let mut array = UInt8Builder::new();
array.append_value(u);
Ok(array.finish().into())
}
fn visit_u16<E>(self, u: u16) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
let mut array = UInt16Builder::new();
array.append_value(u);
Ok(array.finish().into())
}
fn visit_u32<E>(self, u: u32) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
let mut array = UInt32Builder::new();
array.append_value(u);
Ok(array.finish().into())
}
fn visit_u64<E>(self, u: u64) -> Result<Self::Value, E>
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<Field>,
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<A>(self, mut data: A) -> Result<Self::Value, A::Error>
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::<u8>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::UInt16 => {
let mut array = UInt16Builder::new();
while let Some(value) = data.next_element::<u16>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::UInt32 => {
let mut array = UInt32Builder::new();
while let Some(value) = data.next_element::<u32>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::UInt64 => {
let mut array = UInt64Builder::new();
while let Some(value) = data.next_element::<u64>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::Int8 => {
let mut array = Int8Builder::new();
while let Some(value) = data.next_element::<i8>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::Int16 => {
let mut array = Int16Builder::new();
while let Some(value) = data.next_element::<i16>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::Int32 => {
let mut array = Int32Builder::new();
while let Some(value) = data.next_element::<i32>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::Int64 => {
let mut array = Int64Builder::new();
while let Some(value) = data.next_element::<i64>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::Float32 => {
let mut array = Float32Builder::new();
while let Some(value) = data.next_element::<f32>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::Float64 => {
let mut array = Float64Builder::new();
while let Some(value) = data.next_element::<f64>()? {
array.append_value(value);
}
Ok(array.finish().into())
}
DataType::Utf8 => {
let mut array = StringBuilder::new();
while let Some(value) = data.next_element::<String>()? {
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::<Vec<_>>()
.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
}
}
}

+ 111
- 84
python/src/typed/mod.rs View File

@@ -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::<Result<_, _>>()?;

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<String, HashMap<String, Message>>,
) -> eyre::Result<DataType> {
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<String, HashMap<String, Message>>,
) -> eyre::Result<ArrayData> {
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<String, HashMap<String, Message>>,
) -> Result<ArrayData> {
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<u8>).into(),
BasicType::Byte => UInt8Array::from(vec![0u8] as Vec<u8>).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<ArrayData> {
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<String, HashMap<String, Message>>,
) -> Result<ArrayData> {
let defaults = match &m.default.as_deref() {
Some([]) => eyre::bail!("empty default value not supported"),
Some(defaults) => {
let raw_array: Vec<Arc<dyn Array>> = 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::<Result<_, _>>()?;
let default_values = concat(
raw_array
.iter()
.map(|data| data.as_ref())
.collect::<Vec<_>>()
.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)
}

+ 74
- 3
python/src/typed/serialize.rs View File

@@ -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,
},
},


Loading…
Cancel
Save