Skip to main content

azure_storage_blob/stream/
tokio.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4//! Tokio-based stream implementations.
5
6use azure_core::{
7    http::Body,
8    stream::{SeekableStream, DEFAULT_BUFFER_SIZE},
9};
10use std::{
11    future::Future,
12    io::SeekFrom,
13    pin::Pin,
14    sync::Arc,
15    task::{Context, Poll},
16};
17use tokio::{
18    fs::File,
19    io::{AsyncReadExt, AsyncSeekExt},
20    sync::Mutex,
21};
22
23/// Builds a [`FileStream`] from a [`tokio::fs::File`].
24#[derive(Debug)]
25pub struct FileStreamBuilder {
26    file: File,
27    buffer_size: Option<usize>,
28}
29
30impl FileStreamBuilder {
31    fn new(file: File) -> Self {
32        Self {
33            file,
34            buffer_size: None,
35        }
36    }
37
38    /// Sets the size of the buffer to use when reading from the stream.
39    pub fn with_buffer_size(mut self, buffer_size: usize) -> Self {
40        // Not many APIs I looked at use NonZeroUsize which is a bit unwieldy,
41        // but they also don't often protect against this case either.
42        debug_assert!(buffer_size > 0, "buffer_size must be greater than 0");
43
44        self.buffer_size = Some(buffer_size);
45        self
46    }
47
48    /// Builds a [`FileStream`].
49    ///
50    /// # Notes
51    ///
52    /// The [`SeekableStream::len()`] is the file size returned from [`Metadata::len()`](std::fs::Metadata)
53    /// regardless of the initial position of the [`File`].
54    pub async fn build(self) -> azure_core::Result<FileStream> {
55        let file_size = self.file.metadata().await?.len();
56        let buffer_size = self.buffer_size.unwrap_or(DEFAULT_BUFFER_SIZE);
57
58        Ok(FileStream {
59            handle: Arc::new(Mutex::new(self.file)),
60            file_size,
61            buffer_size,
62        })
63    }
64}
65
66/// A stream over a [`tokio::fs::File`] that implements [`SeekableStream`].
67#[derive(Debug, Clone)]
68pub struct FileStream {
69    handle: Arc<Mutex<File>>,
70    file_size: u64,
71    buffer_size: usize,
72}
73
74impl FileStream {
75    /// Creates a new [`FileStreamBuilder`].
76    ///
77    /// # Arguments
78    ///
79    /// * `handle` - An open [`tokio::fs::File`] to stream.
80    ///
81    /// # Notes
82    ///
83    /// `len()` is the file size returned from [`Metadata::len()`](std::fs::Metadata)
84    /// regardless of the initial position of the [`File`].
85    pub fn builder(file: File) -> FileStreamBuilder {
86        FileStreamBuilder::new(file)
87    }
88
89    async fn read(&self, buf: &mut [u8]) -> std::io::Result<usize> {
90        let mut handle = self.handle.lock().await;
91        handle.read(buf).await
92    }
93}
94
95impl From<FileStream> for Body {
96    fn from(stream: FileStream) -> Self {
97        Body::SeekableStream(Box::new(stream))
98    }
99}
100
101#[async_trait::async_trait]
102impl SeekableStream for FileStream {
103    async fn reset(&mut self) -> azure_core::Result<()> {
104        let mut handle = self.handle.lock().await;
105        handle.seek(SeekFrom::Start(0)).await?;
106        Ok(())
107    }
108
109    /// Gets the length of the underlying [`File`].
110    ///
111    /// # Notes
112    ///
113    /// `len()` is the file size returned from [`Metadata::len()`](std::fs::Metadata)
114    /// regardless of the initial position of the [`File`].
115    fn len(&self) -> Option<u64> {
116        Some(self.file_size)
117    }
118
119    fn buffer_size(&self) -> usize {
120        self.buffer_size
121    }
122}
123
124impl futures::io::AsyncRead for FileStream {
125    fn poll_read(
126        self: Pin<&mut Self>,
127        cx: &mut Context<'_>,
128        buf: &mut [u8],
129    ) -> Poll<std::io::Result<usize>> {
130        let fut = self.read(buf);
131        futures::pin_mut!(fut);
132        fut.poll(cx)
133    }
134}
135
136#[cfg(test)]
137mod tests {
138    use super::*;
139    use futures::AsyncReadExt;
140    use std::path::Path;
141
142    async fn open_this_file(buffer_size: Option<usize>) -> FileStream {
143        // file!() returns a workspace-relative path; use a relative traversal
144        // from the crate directory to the workspace root.
145        let path = Path::new(env!("CARGO_MANIFEST_DIR"))
146            .join("../../..")
147            .join(file!());
148        let file = File::open(&path).await.unwrap();
149        let mut builder = FileStream::builder(file);
150        if let Some(size) = buffer_size {
151            builder = builder.with_buffer_size(size);
152        }
153        builder.build().await.unwrap()
154    }
155
156    #[tokio::test]
157    async fn stream_large_chunks() {
158        let mut stream = open_this_file(None).await;
159        let expected_len: usize = stream.len().unwrap().try_into().unwrap();
160        assert!(expected_len > 0);
161
162        let mut buf = vec![0u8; expected_len];
163        let n = stream.read_to_end(&mut buf).await.unwrap();
164        assert_eq!(n, expected_len);
165    }
166
167    #[tokio::test]
168    async fn stream_small_chunks() {
169        const BUFFER_SIZE: usize = 8;
170
171        let stream = open_this_file(Some(BUFFER_SIZE)).await;
172        assert_eq!(stream.buffer_size(), BUFFER_SIZE);
173
174        let expected_len: usize = stream.len().unwrap().try_into().unwrap();
175        let mut total_read = 0;
176        let mut buf = vec![0u8; BUFFER_SIZE];
177        loop {
178            let n = stream.read(&mut buf).await.unwrap();
179            if n == 0 {
180                break;
181            }
182            total_read += n;
183        }
184        assert_eq!(total_read, expected_len);
185    }
186
187    #[tokio::test]
188    async fn reset() {
189        let mut stream = open_this_file(None).await;
190        let expected_len: usize = stream.len().unwrap().try_into().unwrap();
191
192        // First full read.
193        let mut buf1 = vec![0u8; expected_len];
194        let n = stream.read(&mut buf1).await.unwrap();
195        assert_eq!(n, expected_len);
196
197        // Reset and read again.
198        stream.reset().await.unwrap();
199        let mut buf2 = vec![0u8; expected_len];
200        let n = stream.read(&mut buf2).await.unwrap();
201        assert_eq!(n, expected_len);
202        assert_eq!(buf1, buf2);
203    }
204}