1use std::collections::HashMap;
20use std::io::{self, Read, Seek, SeekFrom};
21use std::sync::{Arc, OnceLock, RwLock};
22
23use oxideav_core::{BytesSource, Error, Result};
24
25use crate::uri;
26
27fn table() -> &'static RwLock<HashMap<String, Arc<Vec<u8>>>> {
29 static T: OnceLock<RwLock<HashMap<String, Arc<Vec<u8>>>>> = OnceLock::new();
30 T.get_or_init(|| RwLock::new(HashMap::new()))
31}
32
33pub fn put<I: Into<String>>(id: I, data: Vec<u8>) {
37 table()
38 .write()
39 .expect("mem:// table poisoned")
40 .insert(id.into(), Arc::new(data));
41}
42
43pub fn remove(id: &str) -> bool {
45 table()
46 .write()
47 .expect("mem:// table poisoned")
48 .remove(id)
49 .is_some()
50}
51
52pub fn clear() {
54 table().write().expect("mem:// table poisoned").clear();
55}
56
57pub fn open_mem(uri_str: &str) -> Result<Box<dyn BytesSource>> {
66 let (scheme, rest) = uri::split(uri_str);
67 if scheme != "mem" {
68 return Err(Error::invalid(format!(
69 "mem driver invoked on non-mem URI: {uri_str}"
70 )));
71 }
72 let id = rest;
73 if id.is_empty() {
74 return Err(Error::invalid("mem:// URI requires a non-empty id"));
75 }
76 if id.contains('/') {
77 return Err(Error::invalid(format!(
78 "mem:// id must not contain '/': {id}"
79 )));
80 }
81 let guard = table().read().expect("mem:// table poisoned");
82 let buf = guard
83 .get(id)
84 .ok_or_else(|| Error::invalid(format!("mem:// id '{id}' is not registered")))?;
85 Ok(Box::new(MemReader::new(Arc::clone(buf))))
86}
87
88struct MemReader {
92 buf: Arc<Vec<u8>>,
93 pos: u64,
94}
95
96impl MemReader {
97 fn new(buf: Arc<Vec<u8>>) -> Self {
98 Self { buf, pos: 0 }
99 }
100}
101
102impl Read for MemReader {
103 fn read(&mut self, out: &mut [u8]) -> io::Result<usize> {
104 if out.is_empty() {
105 return Ok(0);
106 }
107 let len = self.buf.len() as u64;
108 if self.pos >= len {
109 return Ok(0);
110 }
111 let avail = (len - self.pos) as usize;
112 let n = out.len().min(avail);
113 let start = self.pos as usize;
114 out[..n].copy_from_slice(&self.buf[start..start + n]);
115 self.pos += n as u64;
116 Ok(n)
117 }
118}
119
120impl Seek for MemReader {
121 fn seek(&mut self, from: SeekFrom) -> io::Result<u64> {
122 let len = self.buf.len() as u64;
123 let new_pos = match from {
124 SeekFrom::Start(n) => n,
125 SeekFrom::End(d) => add_signed(len, d)?,
126 SeekFrom::Current(d) => add_signed(self.pos, d)?,
127 };
128 self.pos = new_pos;
129 Ok(self.pos)
130 }
131}
132
133fn add_signed(base: u64, delta: i64) -> io::Result<u64> {
134 let result = if delta >= 0 {
135 base.checked_add(delta as u64)
136 } else {
137 base.checked_sub(delta.unsigned_abs())
138 };
139 result.ok_or_else(|| {
140 io::Error::new(
141 io::ErrorKind::InvalidInput,
142 "mem:// reader: seek resolves to a negative or overflowing position",
143 )
144 })
145}
146
147#[cfg(test)]
148mod tests {
149 use std::io::{Read, Seek, SeekFrom};
150
151 use super::*;
152
153 fn fresh_id() -> String {
154 use std::sync::atomic::{AtomicU64, Ordering};
155 static N: AtomicU64 = AtomicU64::new(0);
156 format!("test-{}", N.fetch_add(1, Ordering::Relaxed))
157 }
158
159 #[test]
160 fn put_open_read_roundtrip() {
161 let id = fresh_id();
162 put(&id, b"hello, mem://".to_vec());
163 let mut r = open_mem(&format!("mem://{id}")).unwrap();
164 let mut buf = Vec::new();
165 r.read_to_end(&mut buf).unwrap();
166 assert_eq!(buf, b"hello, mem://");
167 assert!(remove(&id));
168 }
169
170 #[test]
171 fn open_supports_seek() {
172 let id = fresh_id();
173 put(&id, (0..=255u8).collect());
174 let mut r = open_mem(&format!("mem://{id}")).unwrap();
175 r.seek(SeekFrom::Start(100)).unwrap();
176 let mut byte = [0u8; 1];
177 r.read_exact(&mut byte).unwrap();
178 assert_eq!(byte[0], 100);
179 let end = r.seek(SeekFrom::End(0)).unwrap();
180 assert_eq!(end, 256);
181 assert!(remove(&id));
182 }
183
184 #[test]
185 fn unknown_id_errors() {
186 let r = open_mem("mem://does-not-exist-xyz");
187 assert!(r.is_err());
188 }
189
190 #[test]
191 fn empty_id_rejected() {
192 let r = open_mem("mem://");
193 assert!(r.is_err());
194 }
195
196 #[test]
197 fn slash_in_id_rejected() {
198 let r = open_mem("mem://foo/bar");
199 assert!(r.is_err());
200 }
201
202 #[test]
203 fn wrong_scheme_rejected() {
204 let r = open_mem("file:///tmp/x");
205 assert!(r.is_err());
206 }
207
208 #[test]
209 fn seek_past_end_then_read_returns_zero() {
210 let id = fresh_id();
211 put(&id, b"abcdef".to_vec());
212 let mut r = open_mem(&format!("mem://{id}")).unwrap();
213 r.seek(SeekFrom::Start(100)).unwrap();
214 let mut buf = [0u8; 8];
215 assert_eq!(r.read(&mut buf).unwrap(), 0);
216 r.seek(SeekFrom::Start(2)).unwrap();
218 let mut chunk = [0u8; 3];
219 r.read_exact(&mut chunk).unwrap();
220 assert_eq!(&chunk, b"cde");
221 assert!(remove(&id));
222 }
223
224 #[test]
225 fn seek_before_zero_errors() {
226 let id = fresh_id();
227 put(&id, b"xy".to_vec());
228 let mut r = open_mem(&format!("mem://{id}")).unwrap();
229 let err = r.seek(SeekFrom::Current(-1));
230 assert!(err.is_err());
231 let err = r.seek(SeekFrom::End(-100));
232 assert!(err.is_err());
233 assert!(remove(&id));
234 }
235
236 #[test]
237 fn large_buffer_open_does_not_copy() {
238 let id = fresh_id();
243 let big: Vec<u8> = (0..(2 * 1024 * 1024)).map(|i| (i & 0xff) as u8).collect();
244 put(&id, big.clone());
245 let mut readers = Vec::with_capacity(16);
248 for _ in 0..16 {
249 readers.push(open_mem(&format!("mem://{id}")).unwrap());
250 }
251 for r in readers.iter_mut() {
253 let mut head = [0u8; 8];
254 r.read_exact(&mut head).unwrap();
255 assert_eq!(head, [0, 1, 2, 3, 4, 5, 6, 7]);
256 }
257 let pos = readers[0].stream_position().unwrap();
259 assert_eq!(pos, 8);
260 assert!(remove(&id));
261 }
262
263 #[test]
264 fn multiple_opens_are_independent() {
265 let id = fresh_id();
266 put(&id, b"AAAAAAAA".to_vec());
267 let mut r1 = open_mem(&format!("mem://{id}")).unwrap();
268 let mut r2 = open_mem(&format!("mem://{id}")).unwrap();
269 let mut a = [0u8; 4];
270 r1.read_exact(&mut a).unwrap();
271 let mut b = [0u8; 1];
273 r2.read_exact(&mut b).unwrap();
274 assert_eq!(b[0], b'A');
275 let pos1 = r1.stream_position().unwrap();
277 let pos2 = r2.stream_position().unwrap();
278 assert_eq!(pos1, 4);
279 assert_eq!(pos2, 1);
280 assert!(remove(&id));
281 }
282}