Skip to main content

mtp_mount/
buffer.rs

1use std::collections::HashMap;
2use std::io;
3use std::io::{Read, Seek, SeekFrom, Write};
4use std::path::{Path, PathBuf};
5
6use crate::error::MountError;
7
8/// Per-file write buffer backed by an unlinked temp file in the spool dir.
9pub struct FileBuffer {
10    pub inode: u64,
11    #[allow(dead_code)]
12    pub original_size: u64,
13    file: std::fs::File,
14    len: u64,
15    dirty: bool,
16}
17
18impl FileBuffer {
19    /// Creates the backing temp file in `spool_dir`. The whole file being
20    /// written buffers here before upload, so `spool_dir` must be disk-backed
21    /// (see [`crate::spool`]).
22    pub fn new(inode: u64, original_size: u64, spool_dir: &Path) -> io::Result<Self> {
23        Ok(Self {
24            inode,
25            original_size,
26            file: tempfile::tempfile_in(spool_dir)?,
27            len: 0,
28            dirty: false,
29        })
30    }
31
32    /// Write `data` at `offset`, growing the file and zero-filling gaps as needed.
33    /// Returns the number of bytes written.
34    pub fn write_at(&mut self, offset: i64, data: &[u8]) -> io::Result<u32> {
35        let offset = offset as u64;
36        let end = offset + data.len() as u64;
37
38        // Zero-fill gaps if writing past current end
39        if offset > self.len {
40            self.file.seek(SeekFrom::Start(self.len))?;
41            let gap = offset - self.len;
42            let zeros = vec![0u8; gap as usize];
43            self.file.write_all(&zeros)?;
44        }
45
46        self.file.seek(SeekFrom::Start(offset))?;
47        self.file.write_all(data)?;
48
49        if end > self.len {
50            self.len = end;
51        }
52        self.dirty = true;
53        Ok(data.len() as u32)
54    }
55
56    /// Read up to `size` bytes starting at `offset`. Returns fewer bytes if
57    /// the offset is near or past the end.
58    pub fn read_at(&mut self, offset: i64, size: u32) -> io::Result<Vec<u8>> {
59        let offset = offset as u64;
60        if offset >= self.len {
61            return Ok(Vec::new());
62        }
63        let available = (self.len - offset).min(size as u64) as usize;
64        let mut buf = vec![0u8; available];
65        self.file.seek(SeekFrom::Start(offset))?;
66        self.file.read_exact(&mut buf)?;
67        Ok(buf)
68    }
69
70    pub fn len(&self) -> u64 {
71        self.len
72    }
73
74    #[allow(dead_code)]
75    pub fn is_empty(&self) -> bool {
76        self.len == 0
77    }
78
79    pub fn is_dirty(&self) -> bool {
80        self.dirty
81    }
82
83    /// Consume the buffer and return the backing file.
84    pub fn into_file(self) -> std::fs::File {
85        self.file
86    }
87}
88
89/// Manages in-progress file writes, mapping file handles to their buffers.
90pub struct WriteBuffer {
91    buffers: HashMap<u64, FileBuffer>,
92    spool_dir: PathBuf,
93}
94
95impl WriteBuffer {
96    /// `spool_dir` is where every buffer's backing file lives; resolve it with
97    /// [`crate::spool::spool_dir_from_env`] and prepare it before calling this.
98    pub fn new(spool_dir: PathBuf) -> Self {
99        Self {
100            buffers: HashMap::new(),
101            spool_dir,
102        }
103    }
104
105    /// Register a new write buffer for the given file handle.
106    pub fn open(&mut self, fh: u64, inode: u64, original_size: u64) -> io::Result<&mut FileBuffer> {
107        use std::collections::hash_map::Entry;
108        match self.buffers.entry(fh) {
109            Entry::Occupied(e) => Ok(e.into_mut()),
110            Entry::Vacant(e) => {
111                let fb = FileBuffer::new(inode, original_size, &self.spool_dir)?;
112                Ok(e.insert(fb))
113            }
114        }
115    }
116
117    /// Write data at offset into the buffer for `fh`.
118    pub fn write(&mut self, fh: u64, offset: i64, data: &[u8]) -> Result<u32, MountError> {
119        let buf = self
120            .buffers
121            .get_mut(&fh)
122            .ok_or_else(|| MountError::Other(format!("no buffer for file handle {fh}")))?;
123        Ok(buf.write_at(offset, data)?)
124    }
125
126    /// Read data from the buffer for `fh`.
127    pub fn read(&mut self, fh: u64, offset: i64, size: u32) -> Result<Vec<u8>, MountError> {
128        let buf = self
129            .buffers
130            .get_mut(&fh)
131            .ok_or_else(|| MountError::Other(format!("no buffer for file handle {fh}")))?;
132        Ok(buf.read_at(offset, size)?)
133    }
134
135    /// Remove and return the buffer for flushing to MTP.
136    pub fn flush(&mut self, fh: u64) -> Option<FileBuffer> {
137        self.buffers.remove(&fh)
138    }
139
140    /// Alias for `flush` -- used at file close time.
141    pub fn close(&mut self, fh: u64) -> Option<FileBuffer> {
142        self.flush(fh)
143    }
144
145    pub fn is_open(&self, fh: u64) -> bool {
146        self.buffers.contains_key(&fh)
147    }
148
149    /// Current buffered size for the given file handle.
150    pub fn size(&self, fh: u64) -> Option<u64> {
151        self.buffers.get(&fh).map(|b| b.len())
152    }
153}
154
155#[cfg(test)]
156mod tests {
157    use super::*;
158
159    /// The buffers are unlinked temp files, so the system temp dir is fine for
160    /// tests; production resolves a disk-backed spool dir instead.
161    fn new_write_buffer() -> WriteBuffer {
162        WriteBuffer::new(std::env::temp_dir())
163    }
164
165    #[test]
166    fn test_new_buffer_empty() {
167        let wb = new_write_buffer();
168        assert!(!wb.is_open(1));
169        assert!(wb.size(1).is_none());
170    }
171
172    #[test]
173    fn test_open_creates_buffer() {
174        let mut wb = new_write_buffer();
175        wb.open(1, 100, 0).unwrap();
176        assert!(wb.is_open(1));
177        assert_eq!(wb.size(1), Some(0));
178    }
179
180    #[test]
181    fn test_write_sequential() {
182        let mut wb = new_write_buffer();
183        wb.open(1, 100, 0).unwrap();
184        wb.write(1, 0, b"hello").unwrap();
185        wb.write(1, 5, b" world").unwrap();
186        assert_eq!(wb.size(1), Some(11));
187        let data = wb.read(1, 0, 11).unwrap();
188        assert_eq!(&data, b"hello world");
189    }
190
191    #[test]
192    fn test_write_at_offset() {
193        let mut wb = new_write_buffer();
194        wb.open(1, 100, 0).unwrap();
195        wb.write(1, 5, b"abc").unwrap();
196        assert_eq!(wb.size(1), Some(8));
197        let data = wb.read(1, 0, 8).unwrap();
198        assert_eq!(&data, b"\0\0\0\0\0abc");
199    }
200
201    #[test]
202    fn test_write_overwrite() {
203        let mut wb = new_write_buffer();
204        wb.open(1, 100, 0).unwrap();
205        wb.write(1, 0, b"hello").unwrap();
206        wb.write(1, 1, b"ELL").unwrap();
207        let data = wb.read(1, 0, 5).unwrap();
208        assert_eq!(&data, b"hELLo");
209    }
210
211    #[test]
212    fn test_read_back() {
213        let mut wb = new_write_buffer();
214        wb.open(1, 100, 0).unwrap();
215        wb.write(1, 0, b"test data").unwrap();
216        let data = wb.read(1, 5, 4).unwrap();
217        assert_eq!(&data, b"data");
218    }
219
220    #[test]
221    fn test_read_past_end() {
222        let mut wb = new_write_buffer();
223        wb.open(1, 100, 0).unwrap();
224        wb.write(1, 0, b"short").unwrap();
225        let data = wb.read(1, 3, 100).unwrap();
226        assert_eq!(&data, b"rt");
227    }
228
229    #[test]
230    fn test_read_empty() {
231        let mut wb = new_write_buffer();
232        wb.open(1, 100, 0).unwrap();
233        let data = wb.read(1, 0, 10).unwrap();
234        assert!(data.is_empty());
235    }
236
237    #[test]
238    fn test_flush_returns_data() {
239        let mut wb = new_write_buffer();
240        wb.open(1, 100, 0).unwrap();
241        wb.write(1, 0, b"flush me").unwrap();
242        let fb = wb.flush(1).unwrap();
243        assert_eq!(fb.inode, 100);
244        let mut file = fb.into_file();
245        file.seek(SeekFrom::Start(0)).unwrap();
246        let mut contents = Vec::new();
247        file.read_to_end(&mut contents).unwrap();
248        assert_eq!(&contents, b"flush me");
249    }
250
251    #[test]
252    fn test_flush_removes_buffer() {
253        let mut wb = new_write_buffer();
254        wb.open(1, 100, 0).unwrap();
255        wb.write(1, 0, b"data").unwrap();
256        wb.flush(1);
257        assert!(!wb.is_open(1));
258    }
259
260    #[test]
261    fn test_dirty_tracking() {
262        let mut wb = new_write_buffer();
263        let fb = wb.open(1, 100, 0).unwrap();
264        assert!(!fb.is_dirty());
265        wb.write(1, 0, b"x").unwrap();
266        // Need to access via flush since we can't borrow after write through wb
267        let fb = wb.flush(1).unwrap();
268        assert!(fb.is_dirty());
269    }
270
271    #[test]
272    fn test_multiple_files() {
273        let mut wb = new_write_buffer();
274        wb.open(1, 100, 0).unwrap();
275        wb.open(2, 200, 0).unwrap();
276        wb.write(1, 0, b"file1").unwrap();
277        wb.write(2, 0, b"file2").unwrap();
278        assert!(wb.is_open(1));
279        assert!(wb.is_open(2));
280        assert_eq!(wb.read(1, 0, 5).unwrap(), b"file1");
281        assert_eq!(wb.read(2, 0, 5).unwrap(), b"file2");
282    }
283
284    #[test]
285    fn test_write_nonexistent_fh() {
286        let mut wb = new_write_buffer();
287        let result = wb.write(999, 0, b"nope");
288        assert!(result.is_err());
289    }
290
291    #[test]
292    fn test_large_write() {
293        let mut wb = new_write_buffer();
294        wb.open(1, 100, 0).unwrap();
295        let big = vec![0xABu8; 2 * 1024 * 1024]; // 2 MB
296        let written = wb.write(1, 0, &big).unwrap();
297        assert_eq!(written, 2 * 1024 * 1024);
298        assert_eq!(wb.size(1), Some(2 * 1024 * 1024));
299    }
300
301    #[test]
302    fn test_sparse_write() {
303        let mut wb = new_write_buffer();
304        wb.open(1, 100, 0).unwrap();
305        wb.write(1, 1000, b"sparse").unwrap();
306        assert_eq!(wb.size(1), Some(1006));
307        let prefix = wb.read(1, 0, 1000).unwrap();
308        assert!(prefix.iter().all(|&b| b == 0));
309        let data = wb.read(1, 1000, 6).unwrap();
310        assert_eq!(&data, b"sparse");
311    }
312}