Skip to main content

mtp_mount/
buffer.rs

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