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