diff --git a/vortex-array/src/dtype/serde/flatbuffers.rs b/vortex-array/src/dtype/serde/flatbuffers.rs index 0f7a16c3d65..8a1e7e6b611 100644 --- a/vortex-array/src/dtype/serde/flatbuffers.rs +++ b/vortex-array/src/dtype/serde/flatbuffers.rs @@ -21,6 +21,7 @@ use vortex_session::VortexSession; use crate::dtype::DType; use crate::dtype::DecimalDType; use crate::dtype::FieldDType; +use crate::dtype::FieldNames; use crate::dtype::MapDType; use crate::dtype::PType; use crate::dtype::StructFields; @@ -69,7 +70,7 @@ impl StructFields { buffer: FlatBuffer, session: VortexSession, ) -> VortexResult { - let names = fb_struct + let names: FieldNames = fb_struct .names() .ok_or_else(|| vortex_err!("failed to parse struct names from flatbuffer"))? .iter() @@ -88,6 +89,14 @@ impl StructFields { }) .collect::>(); + if names.len() != dtypes.len() { + vortex_bail!( + "length mismatch between struct names ({}) and dtypes ({})", + names.len(), + dtypes.len() + ); + } + Ok(StructFields::from_fields(names, dtypes)) } } @@ -556,8 +565,11 @@ impl TryFrom for PType { mod test { use std::sync::Arc; + use flatbuffers::FlatBufferBuilder; use flatbuffers::root; + use vortex_buffer::ByteBuffer; use vortex_flatbuffers::FlatBuffer; + use vortex_flatbuffers::WriteFlatBuffer; use vortex_flatbuffers::WriteFlatBufferExt; use crate::dtype::DType; @@ -742,6 +754,48 @@ mod test { assert_eq!(viewed, eager); } + #[test] + fn test_struct_malformed_flatbuffer() { + let mut fbb = FlatBufferBuilder::new(); + let names = fbb.create_vector::>(&[]); + let dtype_offsets = (0..3) + .map(|_| { + DType::Primitive(PType::I32, Nullability::NonNullable) + .write_flatbuffer(&mut fbb) + .unwrap() + }) + .collect::>(); + let dtypes = fbb.create_vector(&dtype_offsets); + + let struct_table = fb::Struct_::create( + &mut fbb, + &fb::Struct_Args { + names: Some(names), + dtypes: Some(dtypes), + nullable: false, + }, + ); + + let dtype = fb::DType::create( + &mut fbb, + &fb::DTypeArgs { + type_type: fb::Type::Struct_, + type_: Some(struct_table.as_union_value()), + }, + ); + fbb.finish_minimal(dtype); + let (vec, start) = fbb.collapse(); + let end = vec.len(); + let buffer = FlatBuffer::align_from(ByteBuffer::from(vec).slice(start..end)); + + let root_fb = root::(&buffer).unwrap(); + let view = ViewedDType::from_fb_loc(root_fb._tab.loc(), buffer, SESSION.clone()); + + let result = DType::try_from(view); + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("length mismatch")); + } + /// A malformed flatbuffer (here, `dtypes.len() != type_ids.len()`) must round-trip /// to `Err`, not panic. #[test]