1use arrow_array::{ArrayRef, UInt64Array, make_array};
5use arrow_buffer::Buffer;
6use arrow_data::ArrayDataBuilder;
7use arrow_schema::{ArrowError, DataType};
8use arrow_select::take::take;
9
10use crate::DataTypeExt;
11
12type Result<T> = std::result::Result<T, ArrowError>;
13
14pub const INLINE_VALUE_MAX_BYTES: usize = 32;
15
16pub fn extract_scalar_value(array: &ArrayRef, idx: usize) -> Result<ArrayRef> {
17 if idx >= array.len() {
18 return Err(ArrowError::InvalidArgumentError(
19 "Scalar index out of bounds".to_string(),
20 ));
21 }
22
23 take(array.as_ref(), &UInt64Array::from(vec![idx as u64]), None)
24}
25
26fn read_u32(buf: &[u8], offset: &mut usize) -> Result<u32> {
27 if *offset + 4 > buf.len() {
28 return Err(ArrowError::InvalidArgumentError(
29 "Invalid scalar value buffer: unexpected EOF".to_string(),
30 ));
31 }
32 let bytes = [
33 buf[*offset],
34 buf[*offset + 1],
35 buf[*offset + 2],
36 buf[*offset + 3],
37 ];
38 *offset += 4;
39 Ok(u32::from_le_bytes(bytes))
40}
41
42fn read_bytes<'a>(buf: &'a [u8], offset: &mut usize, len: usize) -> Result<&'a [u8]> {
43 if *offset + len > buf.len() {
44 return Err(ArrowError::InvalidArgumentError(
45 "Invalid scalar value buffer: unexpected EOF".to_string(),
46 ));
47 }
48 let slice = &buf[*offset..*offset + len];
49 *offset += len;
50 Ok(slice)
51}
52
53fn write_u32(out: &mut Vec<u8>, v: u32) {
54 out.extend_from_slice(&v.to_le_bytes());
55}
56
57fn write_bytes(out: &mut Vec<u8>, bytes: &[u8]) {
58 out.extend_from_slice(bytes);
59}
60
61pub fn encode_scalar_value_buffer(scalar: &ArrayRef) -> Result<Vec<u8>> {
62 if scalar.len() != 1 || scalar.null_count() != 0 {
63 return Err(ArrowError::InvalidArgumentError(
64 "Scalar value buffer must be a single non-null value".to_string(),
65 ));
66 }
67 let data = scalar.to_data();
68 if data.offset() != 0 {
69 return Err(ArrowError::InvalidArgumentError(
70 "Scalar value buffer must have offset=0".to_string(),
71 ));
72 }
73 if !data.child_data().is_empty() {
74 return Err(ArrowError::InvalidArgumentError(
75 "Scalar value buffer does not support nested types".to_string(),
76 ));
77 }
78
79 let mut out = Vec::with_capacity(128);
86 let buffers = data.buffers();
87 write_u32(&mut out, buffers.len() as u32);
88 for b in buffers {
89 write_u32(&mut out, b.len() as u32);
90 }
91 for b in buffers {
92 write_bytes(&mut out, b.as_slice());
93 }
94 Ok(out)
95}
96
97pub fn decode_scalar_from_value_buffer(
98 data_type: &DataType,
99 value_buffer: &[u8],
100) -> Result<ArrayRef> {
101 if matches!(
102 data_type,
103 DataType::Struct(_) | DataType::FixedSizeList(_, _)
104 ) {
105 return Err(ArrowError::InvalidArgumentError(format!(
106 "Scalar value buffer does not support nested data type {:?}",
107 data_type
108 )));
109 }
110
111 let mut offset = 0;
112 let num_buffers = read_u32(value_buffer, &mut offset)? as usize;
113 let buffer_lens = (0..num_buffers)
114 .map(|_| read_u32(value_buffer, &mut offset).map(|l| l as usize))
115 .collect::<Result<Vec<_>>>()?;
116
117 let mut buffers = Vec::with_capacity(num_buffers);
118 for len in buffer_lens {
119 let bytes = read_bytes(value_buffer, &mut offset, len)?;
120 buffers.push(Buffer::from_vec(bytes.to_vec()));
121 }
122
123 if offset != value_buffer.len() {
124 return Err(ArrowError::InvalidArgumentError(
125 "Invalid scalar value buffer: trailing bytes".to_string(),
126 ));
127 }
128
129 let mut builder = ArrayDataBuilder::new(data_type.clone())
130 .len(1)
131 .null_count(0);
132 for b in buffers {
133 builder = builder.add_buffer(b);
134 }
135 Ok(make_array(builder.build()?))
136}
137
138pub fn decode_scalar_from_inline_value(
139 data_type: &DataType,
140 inline_value: &[u8],
141) -> Result<ArrayRef> {
142 if matches!(data_type, DataType::Boolean) {
149 debug_assert_eq!(
150 inline_value.len(),
151 1,
152 "Invalid boolean inline scalar length (expected 1 byte, got {})",
153 inline_value.len()
154 );
155 } else if let Some(byte_width) = data_type.byte_width_opt() {
156 debug_assert_eq!(
157 inline_value.len(),
158 byte_width,
159 "Inline constant length mismatch for {:?}: expected {} bytes but got {}",
160 data_type,
161 byte_width,
162 inline_value.len()
163 );
164 }
165
166 let data = ArrayDataBuilder::new(data_type.clone())
167 .len(1)
168 .null_count(0)
169 .add_buffer(Buffer::from_vec(inline_value.to_vec()))
170 .build()?;
171 Ok(make_array(data))
172}
173
174pub fn try_inline_value(scalar: &ArrayRef) -> Option<Vec<u8>> {
175 if scalar.null_count() != 0 || scalar.len() != 1 {
176 return None;
177 }
178 let data = scalar.to_data();
179 if !data.child_data().is_empty() {
180 return None;
181 }
182 if data.buffers().len() != 1 {
183 return None;
184 }
185 let bytes = data.buffers()[0].as_slice();
186 if bytes.len() > INLINE_VALUE_MAX_BYTES {
187 return None;
188 }
189 Some(bytes.to_vec())
190}
191
192#[cfg(test)]
193mod tests {
194 use std::sync::Arc;
195
196 use arrow_array::{
197 BooleanArray, DictionaryArray, FixedSizeBinaryArray, Int8Array, Int32Array, StringArray,
198 cast::AsArray, types::Int8Type,
199 };
200
201 use super::*;
202
203 #[test]
204 fn test_extract_scalar_value() {
205 let array: ArrayRef = Arc::new(Int32Array::from(vec![Some(1), None, Some(3)]));
206 let scalar = extract_scalar_value(&array, 2).unwrap();
207 assert_eq!(scalar.len(), 1);
208 assert_eq!(
209 scalar
210 .as_primitive::<arrow_array::types::Int32Type>()
211 .value(0),
212 3
213 );
214 }
215
216 #[test]
217 fn test_extract_scalar_value_from_full_dictionary() {
218 let values = Arc::new(StringArray::from(
219 (0..=i8::MAX)
220 .map(|value| format!("value-{value}"))
221 .collect::<Vec<_>>(),
222 ));
223 let keys = Int8Array::from((0..=i8::MAX).collect::<Vec<_>>());
224 let array: ArrayRef = Arc::new(DictionaryArray::<Int8Type>::new(keys, values));
225
226 let scalar = extract_scalar_value(&array, i8::MAX as usize).unwrap();
227
228 let scalar = scalar.as_dictionary::<Int8Type>();
229 assert_eq!(scalar.len(), 1);
230 assert_eq!(scalar.key(0), Some(i8::MAX as usize));
231 assert_eq!(scalar.values().len(), i8::MAX as usize + 1);
232 }
233
234 #[test]
235 fn test_scalar_value_buffer_utf8_round_trip() {
236 let scalar: ArrayRef = Arc::new(StringArray::from(vec!["hello"]));
237 let buf = encode_scalar_value_buffer(&scalar).unwrap();
238 let decoded = decode_scalar_from_value_buffer(&DataType::Utf8, &buf).unwrap();
239 assert_eq!(decoded.len(), 1);
240 assert_eq!(decoded.null_count(), 0);
241 assert_eq!(decoded.as_string::<i32>().value(0), "hello");
242 }
243
244 #[test]
245 fn test_scalar_value_buffer_fixed_size_binary_round_trip() {
246 let val = vec![0xABu8; 33];
247 let scalar: ArrayRef = Arc::new(
248 FixedSizeBinaryArray::try_from_sparse_iter_with_size(
249 std::iter::once(Some(val.as_slice())),
250 33,
251 )
252 .unwrap(),
253 );
254 let buf = encode_scalar_value_buffer(&scalar).unwrap();
255 let decoded =
256 decode_scalar_from_value_buffer(&DataType::FixedSizeBinary(33), &buf).unwrap();
257 assert_eq!(decoded.len(), 1);
258 assert_eq!(decoded.as_fixed_size_binary().value(0), val.as_slice());
259 }
260
261 #[test]
262 fn test_inline_value_boolean_round_trip() {
263 let scalar: ArrayRef = Arc::new(BooleanArray::from_iter([Some(true)]));
264 let inline = try_inline_value(&scalar).unwrap();
265 let decoded = decode_scalar_from_inline_value(&DataType::Boolean, &inline).unwrap();
266 assert_eq!(decoded.len(), 1);
267 assert_eq!(decoded.null_count(), 0);
268 assert!(decoded.as_boolean().value(0));
269 }
270
271 #[test]
272 fn test_scalar_value_buffer_rejects_nested_type() {
273 let field = Arc::new(arrow_schema::Field::new("item", DataType::Int32, false));
274 let list: ArrayRef = Arc::new(arrow_array::FixedSizeListArray::new(
275 field,
276 2,
277 Arc::new(Int32Array::from(vec![1, 2])),
278 None,
279 ));
280 let scalar = list.slice(0, 1);
281 assert!(encode_scalar_value_buffer(&scalar).is_err());
282 }
283
284 #[test]
285 fn test_decode_scalar_from_value_buffer_rejects_nested_type() {
286 let buf = Vec::<u8>::new();
287 let res =
288 decode_scalar_from_value_buffer(&DataType::Struct(arrow_schema::Fields::empty()), &buf);
289 assert!(res.is_err());
290 }
291
292 #[test]
293 fn test_decode_scalar_from_value_buffer_trailing_bytes() {
294 let mut bytes = Vec::new();
296 bytes.extend_from_slice(&0u32.to_le_bytes());
297 bytes.push(1);
298 let res = decode_scalar_from_value_buffer(&DataType::Int32, &bytes);
299 assert!(res.is_err());
300 }
301}