Skip to main content

lance_io/
utils.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4use std::{cmp::min, num::NonZero, sync::atomic::AtomicU64};
5
6use arrow_array::{
7    ArrayRef,
8    types::{BinaryType, LargeBinaryType, LargeUtf8Type, Utf8Type},
9};
10use arrow_schema::DataType;
11use byteorder::{ByteOrder, LittleEndian};
12use bytes::Bytes;
13use lance_arrow::*;
14use lance_core::deepsize::DeepSizeOf;
15use prost::Message;
16use serde::{Deserialize, Serialize};
17
18use crate::{ReadBatchParams, traits::Reader};
19use crate::{
20    encodings::{AsyncIndex, Decoder, binary::BinaryDecoder, plain::PlainDecoder},
21    traits::ProtoStruct,
22};
23use lance_core::{Error, Result};
24
25pub mod tracking_store;
26
27/// Read a binary array from a [Reader].
28///
29pub async fn read_binary_array(
30    reader: &dyn Reader,
31    data_type: &DataType,
32    nullable: bool,
33    position: usize,
34    length: usize,
35    params: impl Into<ReadBatchParams>,
36) -> Result<ArrayRef> {
37    use arrow_schema::DataType::*;
38    let decoder: Box<dyn Decoder<Output = Result<ArrayRef>> + Send> = match data_type {
39        Utf8 => Box::new(BinaryDecoder::<Utf8Type>::new(
40            reader, position, length, nullable,
41        )),
42        Binary => Box::new(BinaryDecoder::<BinaryType>::new(
43            reader, position, length, nullable,
44        )),
45        LargeUtf8 => Box::new(BinaryDecoder::<LargeUtf8Type>::new(
46            reader, position, length, nullable,
47        )),
48        LargeBinary => Box::new(BinaryDecoder::<LargeBinaryType>::new(
49            reader, position, length, nullable,
50        )),
51        _ => {
52            return Err(Error::invalid_input(format!(
53                "Unsupported binary type: {}",
54                data_type
55            )));
56        }
57    };
58    let fut = decoder.as_ref().get(params.into());
59    fut.await
60}
61
62/// Read a fixed stride array from disk.
63///
64pub async fn read_fixed_stride_array(
65    reader: &dyn Reader,
66    data_type: &DataType,
67    position: usize,
68    length: usize,
69    params: impl Into<ReadBatchParams>,
70) -> Result<ArrayRef> {
71    if !data_type.is_fixed_stride() {
72        return Err(Error::schema(format!(
73            "{data_type} is not a fixed stride type"
74        )));
75    }
76    // TODO: support more than plain encoding here.
77    let decoder = PlainDecoder::new(reader, data_type, position, length)?;
78    decoder.get(params.into()).await
79}
80
81/// Read a protobuf message at file position 'pos'.
82///
83/// We write protobuf by first writing the length of the message as a u32,
84/// followed by the message itself.
85pub async fn read_message<M: Message + Default>(reader: &dyn Reader, pos: usize) -> Result<M> {
86    let file_size = reader.size().await?;
87    // A message is a u32 length prefix followed by its body; both must lie before
88    // the end. A `pos` too close to the end means the reader size is too small
89    // (e.g. a stale cached size). Reject it rather than slice a short buffer and
90    // panic.
91    if pos + 4 > file_size {
92        return Err(Error::io("file size is too small".to_string()));
93    }
94
95    let range = pos..min(pos + reader.block_size(), file_size);
96    let buf = reader.get_range(range.clone()).await?;
97    let msg_len = LittleEndian::read_u32(&buf) as usize;
98
99    if msg_len + 4 > buf.len() {
100        let remaining_range = range.end..min(4 + pos + msg_len, file_size);
101        let remaining_bytes = reader.get_range(remaining_range).await?;
102        let buf = [buf, remaining_bytes].concat();
103        if buf.len() < msg_len + 4 {
104            return Err(Error::io("file size is too small".to_string()));
105        }
106        Ok(M::decode(&buf[4..4 + msg_len])?)
107    } else {
108        Ok(M::decode(&buf[4..4 + msg_len])?)
109    }
110}
111
112/// Read a Protobuf-backed struct at file position: `pos`.
113// TODO: pub(crate)
114pub async fn read_struct<
115    M: Message + Default + 'static,
116    T: ProtoStruct<Proto = M> + TryFrom<M, Error = Error>,
117>(
118    reader: &dyn Reader,
119    pos: usize,
120) -> Result<T> {
121    let msg = read_message::<M>(reader, pos).await?;
122    T::try_from(msg)
123}
124
125pub async fn read_last_block(reader: &dyn Reader) -> object_store::Result<Bytes> {
126    let file_size = reader.size().await?;
127    let block_size = reader.block_size();
128    let begin = file_size.saturating_sub(block_size);
129    reader.get_range(begin..file_size).await
130}
131
132pub fn read_metadata_offset(bytes: &Bytes) -> Result<usize> {
133    let len = bytes.len();
134    if len < 16 {
135        return Err(Error::io(format!(
136            "does not have sufficient data, len: {}, bytes: {:?}",
137            len, bytes
138        )));
139    }
140    let offset_bytes = bytes.slice(len - 16..len - 8);
141    Ok(LittleEndian::read_u64(offset_bytes.as_ref()) as usize)
142}
143
144/// Read the version from the footer bytes
145pub fn read_version(bytes: &Bytes) -> Result<(u16, u16)> {
146    let len = bytes.len();
147    if len < 8 {
148        return Err(Error::io(format!(
149            "does not have sufficient data, len: {}, bytes: {:?}",
150            len, bytes
151        )));
152    }
153
154    let major_version = LittleEndian::read_u16(bytes.slice(len - 8..len - 6).as_ref());
155    let minor_version = LittleEndian::read_u16(bytes.slice(len - 6..len - 4).as_ref());
156    Ok((major_version, minor_version))
157}
158
159/// Read protobuf from a buffer.
160pub fn read_message_from_buf<M: Message + Default>(buf: &Bytes) -> Result<M> {
161    let msg_len = LittleEndian::read_u32(buf) as usize;
162    Ok(M::decode(&buf[4..4 + msg_len])?)
163}
164
165/// Read a Protobuf-backed struct from a buffer.
166pub fn read_struct_from_buf<
167    M: Message + Default,
168    T: ProtoStruct<Proto = M> + TryFrom<M, Error = Error>,
169>(
170    buf: &Bytes,
171) -> Result<T> {
172    let msg: M = read_message_from_buf(buf)?;
173    T::try_from(msg)
174}
175
176/// A cached file size.
177///
178/// This wraps an atomic u64 to allow setting the cached file size without
179/// needed a mutable reference.
180///
181/// Zero is interpreted as unknown.
182#[derive(Debug, DeepSizeOf)]
183pub struct CachedFileSize(AtomicU64);
184
185impl<'de> Deserialize<'de> for CachedFileSize {
186    fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
187    where
188        D: serde::Deserializer<'de>,
189    {
190        let size = Option::<u64>::deserialize(deserializer)?.unwrap_or(0);
191        Ok(Self::new(size))
192    }
193}
194
195impl Serialize for CachedFileSize {
196    fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
197    where
198        S: serde::Serializer,
199    {
200        let size = self.0.load(std::sync::atomic::Ordering::Relaxed);
201        if size == 0 {
202            serializer.serialize_none()
203        } else {
204            serializer.serialize_u64(size)
205        }
206    }
207}
208
209impl From<Option<NonZero<u64>>> for CachedFileSize {
210    fn from(size: Option<NonZero<u64>>) -> Self {
211        match size {
212            Some(size) => Self(AtomicU64::new(size.into())),
213            None => Self(AtomicU64::new(0)),
214        }
215    }
216}
217
218impl Default for CachedFileSize {
219    fn default() -> Self {
220        Self(AtomicU64::new(0))
221    }
222}
223
224impl Clone for CachedFileSize {
225    fn clone(&self) -> Self {
226        Self(AtomicU64::new(
227            self.0.load(std::sync::atomic::Ordering::Relaxed),
228        ))
229    }
230}
231
232impl PartialEq for CachedFileSize {
233    fn eq(&self, other: &Self) -> bool {
234        self.0.load(std::sync::atomic::Ordering::Relaxed)
235            == other.0.load(std::sync::atomic::Ordering::Relaxed)
236    }
237}
238
239impl Eq for CachedFileSize {}
240
241impl CachedFileSize {
242    /// Create a `CachedFileSize` from a raw byte count.
243    ///
244    /// Passing `0` is equivalent to calling [`unknown`](Self::unknown): the
245    /// type interprets zero as "size not yet known".
246    pub fn new(size: u64) -> Self {
247        Self(AtomicU64::new(size))
248    }
249
250    pub fn unknown() -> Self {
251        Self(AtomicU64::new(0))
252    }
253
254    pub fn get(&self) -> Option<NonZero<u64>> {
255        NonZero::new(self.0.load(std::sync::atomic::Ordering::Relaxed))
256    }
257
258    pub fn set(&self, size: NonZero<u64>) {
259        self.0
260            .store(size.into(), std::sync::atomic::Ordering::Relaxed);
261    }
262}
263
264#[cfg(test)]
265mod tests {
266    use bytes::Bytes;
267    use object_store::path::Path;
268
269    use crate::{
270        Error, Result,
271        object_reader::CloudObjectReader,
272        object_store::{DEFAULT_DOWNLOAD_RETRY_COUNT, ObjectStore},
273        object_writer::ObjectWriter,
274        traits::{ProtoStruct, WriteExt, Writer},
275        utils::read_struct,
276    };
277
278    // Bytes is a prost::Message, since we don't have any .proto files in this crate we
279    // can use it to simulate a real message object.
280    #[derive(Debug, PartialEq)]
281    struct BytesWrapper(Bytes);
282
283    impl ProtoStruct for BytesWrapper {
284        type Proto = Bytes;
285    }
286
287    impl From<&BytesWrapper> for Bytes {
288        fn from(value: &BytesWrapper) -> Self {
289            value.0.clone()
290        }
291    }
292
293    impl TryFrom<Bytes> for BytesWrapper {
294        type Error = Error;
295        fn try_from(value: Bytes) -> Result<Self> {
296            Ok(Self(value))
297        }
298    }
299
300    #[tokio::test]
301    async fn test_write_proto_structs() {
302        let store = ObjectStore::memory();
303        let path = Path::from("/foo");
304
305        let mut object_writer = ObjectWriter::new(&store, &path).await.unwrap();
306        assert_eq!(object_writer.tell().await.unwrap(), 0);
307
308        let some_message = BytesWrapper(Bytes::from(vec![10, 20, 30]));
309
310        let pos = object_writer.write_struct(&some_message).await.unwrap();
311        assert_eq!(pos, 0);
312        object_writer.shutdown().await.unwrap();
313
314        let object_reader =
315            CloudObjectReader::new(store.inner, path, 1024, None, DEFAULT_DOWNLOAD_RETRY_COUNT)
316                .unwrap();
317        let actual: BytesWrapper = read_struct(&object_reader, pos).await.unwrap();
318        assert_eq!(some_message, actual);
319    }
320
321    #[tokio::test]
322    async fn test_copy_reader_to_writer() {
323        let store = ObjectStore::memory();
324        let src = Path::from("/src");
325        let dst = Path::from("/dst");
326        store.put(&src, b"abcdef").await.unwrap();
327
328        let reader = store.open(&src).await.unwrap();
329        let mut writer = store.create(&dst).await.unwrap();
330        let copied = writer.copy_from_reader(reader.as_ref()).await.unwrap();
331        writer.shutdown().await.unwrap();
332
333        assert_eq!(copied, 6);
334        assert_eq!(store.read_one_all(&dst).await.unwrap().as_ref(), b"abcdef");
335    }
336
337    #[tokio::test]
338    async fn test_copy_reader_range_to_writer() {
339        let store = ObjectStore::memory();
340        let src = Path::from("/src-range");
341        let dst = Path::from("/dst-range");
342        store.put(&src, b"abcdef").await.unwrap();
343
344        let reader = store.open(&src).await.unwrap();
345        let mut writer = store.create(&dst).await.unwrap();
346        let copied = writer
347            .copy_range_from_reader(reader.as_ref(), 2..5)
348            .await
349            .unwrap();
350        writer.shutdown().await.unwrap();
351
352        assert_eq!(copied, 3);
353        assert_eq!(store.read_one_all(&dst).await.unwrap().as_ref(), b"cde");
354    }
355}