1use std::{cmp::min, num::NonZero, ops::Range, sync::atomic::AtomicU64};
5
6use byteorder::{ByteOrder, LittleEndian};
7use bytes::{Bytes, BytesMut};
8use futures::{Stream, StreamExt, TryStreamExt};
9use lance_core::deepsize::DeepSizeOf;
10use prost::Message;
11use serde::{Deserialize, Serialize};
12
13use crate::traits::{ProtoStruct, Reader};
14use lance_core::{Error, Result};
15
16pub mod tracking_store;
17
18pub const METADATA_READ_CHUNK_SIZE: usize = 16 * 1024 * 1024;
28
29pub fn read_range_in_chunks(
35 reader: &dyn Reader,
36 range: Range<usize>,
37 chunk_size: usize,
38) -> impl Stream<Item = object_store::Result<Bytes>> + '_ {
39 let end = range.end;
40 let chunk_ranges = range
41 .step_by(chunk_size)
42 .map(move |start| start..min(start + chunk_size, end));
43 futures::stream::iter(chunk_ranges.map(|chunk| reader.get_range(chunk)))
44 .buffered(reader.io_parallelism().max(1))
45}
46
47pub async fn read_message<M: Message + Default>(reader: &dyn Reader, pos: usize) -> Result<M> {
52 let file_size = reader.size().await?;
53 if pos + 4 > file_size {
58 return Err(Error::io("file size is too small".to_string()));
59 }
60
61 let range = pos..min(pos + reader.block_size(), file_size);
62 let buf = reader.get_range(range.clone()).await?;
63 let msg_len = LittleEndian::read_u32(&buf) as usize;
64
65 if msg_len + 4 > buf.len() {
66 let remaining_range = range.end..min(4 + pos + msg_len, file_size);
67 let mut full = BytesMut::with_capacity(buf.len() + remaining_range.len());
71 full.extend_from_slice(&buf);
72 let mut chunks = read_range_in_chunks(reader, remaining_range, METADATA_READ_CHUNK_SIZE);
73 while let Some(chunk) = chunks.try_next().await? {
74 full.extend_from_slice(&chunk);
75 }
76 if full.len() < msg_len + 4 {
77 return Err(Error::io("file size is too small".to_string()));
78 }
79 Ok(M::decode(&full[4..4 + msg_len])?)
80 } else {
81 Ok(M::decode(&buf[4..4 + msg_len])?)
82 }
83}
84
85pub async fn read_struct<
88 M: Message + Default + 'static,
89 T: ProtoStruct<Proto = M> + TryFrom<M, Error = Error>,
90>(
91 reader: &dyn Reader,
92 pos: usize,
93) -> Result<T> {
94 let msg = read_message::<M>(reader, pos).await?;
95 T::try_from(msg)
96}
97
98pub async fn read_last_block(reader: &dyn Reader) -> object_store::Result<Bytes> {
99 let file_size = reader.size().await?;
100 let block_size = reader.block_size();
101 let begin = file_size.saturating_sub(block_size);
102 reader.get_range(begin..file_size).await
103}
104
105pub fn read_metadata_offset(bytes: &Bytes) -> Result<usize> {
106 let len = bytes.len();
107 if len < 16 {
108 return Err(Error::io(format!(
109 "does not have sufficient data, len: {}, bytes: {:?}",
110 len, bytes
111 )));
112 }
113 let offset_bytes = bytes.slice(len - 16..len - 8);
114 Ok(LittleEndian::read_u64(offset_bytes.as_ref()) as usize)
115}
116
117pub fn read_version(bytes: &Bytes) -> Result<(u16, u16)> {
119 let len = bytes.len();
120 if len < 8 {
121 return Err(Error::io(format!(
122 "does not have sufficient data, len: {}, bytes: {:?}",
123 len, bytes
124 )));
125 }
126
127 let major_version = LittleEndian::read_u16(bytes.slice(len - 8..len - 6).as_ref());
128 let minor_version = LittleEndian::read_u16(bytes.slice(len - 6..len - 4).as_ref());
129 Ok((major_version, minor_version))
130}
131
132pub fn read_message_from_buf<M: Message + Default>(buf: &Bytes) -> Result<M> {
134 let msg_len = LittleEndian::read_u32(buf) as usize;
135 Ok(M::decode(&buf[4..4 + msg_len])?)
136}
137
138pub fn read_struct_from_buf<
140 M: Message + Default,
141 T: ProtoStruct<Proto = M> + TryFrom<M, Error = Error>,
142>(
143 buf: &Bytes,
144) -> Result<T> {
145 let msg: M = read_message_from_buf(buf)?;
146 T::try_from(msg)
147}
148
149#[derive(Debug, DeepSizeOf)]
156pub struct CachedFileSize(AtomicU64);
157
158impl<'de> Deserialize<'de> for CachedFileSize {
159 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
160 where
161 D: serde::Deserializer<'de>,
162 {
163 let size = Option::<u64>::deserialize(deserializer)?.unwrap_or(0);
164 Ok(Self::new(size))
165 }
166}
167
168impl Serialize for CachedFileSize {
169 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
170 where
171 S: serde::Serializer,
172 {
173 let size = self.0.load(std::sync::atomic::Ordering::Relaxed);
174 if size == 0 {
175 serializer.serialize_none()
176 } else {
177 serializer.serialize_u64(size)
178 }
179 }
180}
181
182impl From<Option<NonZero<u64>>> for CachedFileSize {
183 fn from(size: Option<NonZero<u64>>) -> Self {
184 match size {
185 Some(size) => Self(AtomicU64::new(size.into())),
186 None => Self(AtomicU64::new(0)),
187 }
188 }
189}
190
191impl Default for CachedFileSize {
192 fn default() -> Self {
193 Self(AtomicU64::new(0))
194 }
195}
196
197impl Clone for CachedFileSize {
198 fn clone(&self) -> Self {
199 Self(AtomicU64::new(
200 self.0.load(std::sync::atomic::Ordering::Relaxed),
201 ))
202 }
203}
204
205impl PartialEq for CachedFileSize {
206 fn eq(&self, other: &Self) -> bool {
207 self.0.load(std::sync::atomic::Ordering::Relaxed)
208 == other.0.load(std::sync::atomic::Ordering::Relaxed)
209 }
210}
211
212impl Eq for CachedFileSize {}
213
214impl CachedFileSize {
215 pub fn new(size: u64) -> Self {
220 Self(AtomicU64::new(size))
221 }
222
223 pub fn unknown() -> Self {
224 Self(AtomicU64::new(0))
225 }
226
227 pub fn get(&self) -> Option<NonZero<u64>> {
228 NonZero::new(self.0.load(std::sync::atomic::Ordering::Relaxed))
229 }
230
231 pub fn set(&self, size: NonZero<u64>) {
232 self.0
233 .store(size.into(), std::sync::atomic::Ordering::Relaxed);
234 }
235}
236
237#[cfg(test)]
238mod tests {
239 use bytes::{Bytes, BytesMut};
240 use futures::TryStreamExt;
241 use object_store::path::Path;
242
243 use crate::{
244 Error, Result,
245 object_reader::CloudObjectReader,
246 object_store::{DEFAULT_DOWNLOAD_RETRY_COUNT, ObjectStore},
247 object_writer::ObjectWriter,
248 traits::{ProtoStruct, WriteExt, Writer},
249 utils::{METADATA_READ_CHUNK_SIZE, read_range_in_chunks, read_struct},
250 };
251
252 #[derive(Debug, PartialEq)]
255 struct BytesWrapper(Bytes);
256
257 impl ProtoStruct for BytesWrapper {
258 type Proto = Bytes;
259 }
260
261 impl From<&BytesWrapper> for Bytes {
262 fn from(value: &BytesWrapper) -> Self {
263 value.0.clone()
264 }
265 }
266
267 impl TryFrom<Bytes> for BytesWrapper {
268 type Error = Error;
269 fn try_from(value: Bytes) -> Result<Self> {
270 Ok(Self(value))
271 }
272 }
273
274 #[tokio::test]
275 async fn test_write_proto_structs() {
276 let store = ObjectStore::memory();
277 let path = Path::from("/foo");
278
279 let mut object_writer = ObjectWriter::new(&store, &path).await.unwrap();
280 assert_eq!(object_writer.tell().await.unwrap(), 0);
281
282 let some_message = BytesWrapper(Bytes::from(vec![10, 20, 30]));
283
284 let pos = object_writer.write_struct(&some_message).await.unwrap();
285 assert_eq!(pos, 0);
286 object_writer.shutdown().await.unwrap();
287
288 let object_reader =
289 CloudObjectReader::new(store.inner, path, 1024, None, DEFAULT_DOWNLOAD_RETRY_COUNT)
290 .unwrap();
291 let actual: BytesWrapper = read_struct(&object_reader, pos).await.unwrap();
292 assert_eq!(some_message, actual);
293 }
294
295 #[tokio::test]
296 async fn test_read_range_in_chunks_reassembles_in_order() {
297 let store = ObjectStore::memory();
298 let path = Path::from("/chunked");
299 let data: Vec<u8> = (0..10 * 1024 + 37).map(|i| (i % 251) as u8).collect();
302 store.put(&path, &data).await.unwrap();
303 let reader = store.open(&path).await.unwrap();
304
305 let range = 5..data.len() - 3;
306 let mut assembled = BytesMut::new();
307 let mut chunks = read_range_in_chunks(reader.as_ref(), range.clone(), 1024);
308 while let Some(chunk) = chunks.try_next().await.unwrap() {
309 assembled.extend_from_slice(&chunk);
310 }
311 assert_eq!(assembled.as_ref(), &data[range]);
312 }
313
314 #[tokio::test]
315 async fn test_read_range_in_chunks_zero_parallelism_reader() {
316 let store = ObjectStore::memory();
319 let path = Path::from("/zero_parallelism");
320 let data: Vec<u8> = (0..4096).map(|i| (i % 249) as u8).collect();
321 store.put(&path, &data).await.unwrap();
322 let reader =
323 CloudObjectReader::new(store.inner, path, 1024, None, DEFAULT_DOWNLOAD_RETRY_COUNT)
324 .unwrap()
325 .with_io_parallelism(0);
326
327 let assembled = tokio::time::timeout(std::time::Duration::from_secs(5), async {
328 let mut buf = BytesMut::new();
329 let mut chunks = read_range_in_chunks(&reader, 0..data.len(), 1024);
330 while let Some(chunk) = chunks.try_next().await.unwrap() {
331 buf.extend_from_slice(&chunk);
332 }
333 buf
334 })
335 .await
336 .expect("chunked read with a zero-parallelism reader must not hang");
337 assert_eq!(assembled.as_ref(), &data[..]);
338 }
339
340 #[tokio::test]
341 async fn test_read_message_larger_than_chunk_size() {
342 let store = ObjectStore::memory();
345 let path = Path::from("/large_message");
346
347 let mut object_writer = ObjectWriter::new(&store, &path).await.unwrap();
348 let payload: Vec<u8> = (0..METADATA_READ_CHUNK_SIZE + 5 * 1024 * 1024)
349 .map(|i| (i % 253) as u8)
350 .collect();
351 let message = BytesWrapper(Bytes::from(payload));
352 let pos = object_writer.write_struct(&message).await.unwrap();
353 object_writer.shutdown().await.unwrap();
354
355 let object_reader =
356 CloudObjectReader::new(store.inner, path, 4096, None, DEFAULT_DOWNLOAD_RETRY_COUNT)
357 .unwrap();
358 let actual: BytesWrapper = read_struct(&object_reader, pos).await.unwrap();
359 assert_eq!(message, actual);
360 }
361
362 #[tokio::test]
363 async fn test_copy_reader_to_writer() {
364 let store = ObjectStore::memory();
365 let src = Path::from("/src");
366 let dst = Path::from("/dst");
367 store.put(&src, b"abcdef").await.unwrap();
368
369 let reader = store.open(&src).await.unwrap();
370 let mut writer = store.create(&dst).await.unwrap();
371 let copied = writer.copy_from_reader(reader.as_ref()).await.unwrap();
372 writer.shutdown().await.unwrap();
373
374 assert_eq!(copied, 6);
375 assert_eq!(store.read_one_all(&dst).await.unwrap().as_ref(), b"abcdef");
376 }
377
378 #[tokio::test]
379 async fn test_copy_reader_range_to_writer() {
380 let store = ObjectStore::memory();
381 let src = Path::from("/src-range");
382 let dst = Path::from("/dst-range");
383 store.put(&src, b"abcdef").await.unwrap();
384
385 let reader = store.open(&src).await.unwrap();
386 let mut writer = store.create(&dst).await.unwrap();
387 let copied = writer
388 .copy_range_from_reader(reader.as_ref(), 2..5)
389 .await
390 .unwrap();
391 writer.shutdown().await.unwrap();
392
393 assert_eq!(copied, 3);
394 assert_eq!(store.read_one_all(&dst).await.unwrap().as_ref(), b"cde");
395 }
396}