Skip to main content

lance_arrow/
scalar.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4use 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    // Minimal format (RFC): store the Arrow value buffers for a length-1 array.
80    // Null bitmap and child data are intentionally not supported here.
81    //
82    // | u32 num_buffers |
83    // | u32 buffer_0_len | ... | u32 buffer_{n-1}_len |
84    // | buffer_0 bytes | ... | buffer_{n-1} bytes |
85    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    // I expect our input to be safe here, but I added some debug_assert_eq statements just in case.
143    // If they are triggered, we may need to change them to return actual errors.
144    //
145    // Boolean values are bit-packed in Arrow and therefore are not "fixed-stride" in bytes.
146    // As a result, `byte_width_opt()` returns `None` for `DataType::Boolean`, even though a
147    // length-1 scalar can be represented inline using a single byte (matching `try_inline_value`).
148    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        // num_buffers = 0, plus an extra byte
295        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}