Skip to main content

diskann_disk/search/provider/aligned_file_reader/reader/
storage_provider.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6use std::io::Read;
7
8use diskann::ANNResult;
9use diskann_providers::storage::StorageReadProvider;
10use tracing::info;
11
12use crate::search::provider::aligned_file_reader::{traits::AlignedFileReader, AlignedRead, A1};
13
14pub struct StorageProviderAlignedFileReader {
15    data: Vec<u8>,
16}
17
18impl StorageProviderAlignedFileReader {
19    pub fn new(
20        storage_provider: &impl StorageReadProvider,
21        file_name: &str,
22    ) -> ANNResult<StorageProviderAlignedFileReader> {
23        info!("Loading data from {}", file_name);
24        let file_length = storage_provider.get_length(file_name)?;
25
26        let mut data = vec![0u8; file_length as usize];
27        storage_provider
28            .open_reader(file_name)?
29            .read_exact(&mut data)?;
30
31        Ok(StorageProviderAlignedFileReader { data })
32    }
33}
34
35impl AlignedFileReader for StorageProviderAlignedFileReader {
36    type Alignment = A1;
37
38    fn read(&mut self, read_requests: &mut [AlignedRead<u8, A1>]) -> ANNResult<()> {
39        for read in read_requests {
40            let offset = read.offset();
41            let len = read.aligned_buf().len();
42            let aligned_buf = read.aligned_buf_mut();
43            aligned_buf.copy_from_slice(&self.data[offset as usize..offset as usize + len]);
44        }
45
46        Ok(())
47    }
48}
49
50#[cfg(test)]
51mod tests {
52    use std::io::{Seek, SeekFrom};
53
54    use diskann_providers::storage::VirtualStorageProvider;
55    use diskann_utils::test_data_root;
56
57    use super::*;
58    use diskann_quantization::alloc::{AlignedAllocator, Poly};
59
60    fn test_index_path() -> String {
61        "/disk_index_misc/disk_index_siftsmall_learn_256pts_R4_L50_A1.2_aligned_reader_test.index"
62            .to_string()
63    }
64
65    fn setup_reader() -> StorageProviderAlignedFileReader {
66        let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
67        StorageProviderAlignedFileReader::new(&storage_provider, &test_index_path()).unwrap()
68    }
69
70    #[test]
71    fn test_new_aligned_file_reader() {
72        let reader = setup_reader();
73        assert!(!(reader.data.is_empty()));
74    }
75
76    #[test]
77    fn test_read() {
78        let mut reader = setup_reader();
79
80        let read_length = 512;
81        let num_read = 10;
82        let mut aligned_mem =
83            Poly::broadcast(0u8, read_length * num_read, AlignedAllocator::A512).unwrap();
84
85        // create and add AlignedReads to the vector
86        let mut mem_slices: Vec<&mut [u8]> = aligned_mem.chunks_mut(read_length).collect();
87
88        let mut aligned_reads: Vec<AlignedRead<'_, u8>> = mem_slices
89            .iter_mut()
90            .enumerate()
91            .map(|(i, slice)| {
92                let offset = (i * read_length) as u64;
93                AlignedRead::new(offset, slice).unwrap()
94            })
95            .collect();
96
97        let result = reader.read(&mut aligned_reads);
98        assert!(result.is_ok());
99
100        // Assert that the actual data is correct.
101        let file_system = VirtualStorageProvider::new_overlay(test_data_root());
102        let mut file = file_system.open_reader(&test_index_path()).unwrap();
103        for current_read in aligned_reads {
104            let offset = current_read.offset();
105            let mut expected = vec![0; current_read.aligned_buf().len()];
106            file.seek(SeekFrom::Start(offset)).unwrap();
107            file.read_exact(&mut expected).unwrap();
108
109            assert_eq!(
110                expected,
111                current_read.aligned_buf(),
112                "aligned_buf did not contain the expected data"
113            );
114        }
115    }
116}