1use std::fs::File;
22use std::io;
23use std::path::Path;
24
25pub const PAGE_SIZE: usize = 4096;
29
30pub const DEFAULT_BUFFER_SIZE: usize = 65536;
32
33#[derive(Debug)]
36pub enum IoError {
37 MisalignedOffset { offset: u64, required: u64 },
38 MisalignedBufferSize { size: usize, required: usize },
39 FileOpenError(io::Error),
40 IoError(io::Error),
41 LockError(String),
42 InvalidState(String),
43}
44
45impl From<io::Error> for IoError {
46 fn from(err: io::Error) -> Self {
47 IoError::IoError(err)
48 }
49}
50
51#[repr(C, align(4096))]
58pub struct DmaBuffer<const N: usize> {
59 data: [u8; N],
60}
61
62impl<const N: usize> DmaBuffer<N> {
63 pub fn new() -> Self {
69 assert!(
70 N % PAGE_SIZE == 0,
71 "DMA buffer size must be a multiple of PAGE_SIZE (4096), got {}",
72 N
73 );
74
75 Self { data: [0u8; N] }
76 }
77
78 pub fn as_slice(&self) -> &[u8] {
80 &self.data
81 }
82
83 pub fn as_mut_slice(&mut self) -> &mut [u8] {
85 &mut self.data
86 }
87
88 pub fn len(&self) -> usize {
90 N
91 }
92
93 pub fn is_empty(&self) -> bool {
95 N == 0
96 }
97
98 pub fn f64_capacity(&self) -> usize {
100 N / 8
101 }
102}
103
104impl<const N: usize> Default for DmaBuffer<N> {
105 fn default() -> Self {
106 Self::new()
107 }
108}
109
110pub trait ZeroCopyStreamer: Send {
118 fn async_read_chunk(&mut self, offset: u64) -> Result<(), IoError>;
125
126 fn poll_completion(&mut self) -> Option<&[u8]>;
134
135 fn get_active_buffer(&self) -> &[u8];
139
140 fn buffer_size(&self) -> usize;
142}
143
144#[cfg(target_os = "windows")]
151pub struct IocpGridManager {
152 handle: *mut crate::directml_bridge::IocpHandle,
153 buffer_a: DmaBuffer<DEFAULT_BUFFER_SIZE>,
154 buffer_b: DmaBuffer<DEFAULT_BUFFER_SIZE>,
155 active_buffer: BufferId,
156 pending_read: bool,
157}
158
159#[cfg(target_os = "windows")]
160#[derive(Clone, Copy, PartialEq, Eq)]
161enum BufferId {
162 A,
163 B,
164}
165
166#[cfg(target_os = "windows")]
167unsafe impl Send for IocpGridManager {}
168
169#[cfg(target_os = "windows")]
170impl IocpGridManager {
171 pub fn new(file_path: &Path) -> Result<Self, IoError> {
173 let path_str = file_path
174 .to_str()
175 .ok_or_else(|| IoError::InvalidState("Invalid UTF-8 path".to_string()))?;
176
177 unsafe {
178 let mut handle = std::ptr::null_mut();
179 let status = crate::directml_bridge::iocp_create_ffi(
180 path_str.as_ptr(),
181 path_str.len(),
182 &mut handle,
183 );
184
185 if status != crate::directml_bridge::DmlStatus::Success {
186 return Err(IoError::IoError(io::Error::new(
187 io::ErrorKind::Other,
188 status.message(),
189 )));
190 }
191
192 Ok(Self {
193 handle,
194 buffer_a: DmaBuffer::new(),
195 buffer_b: DmaBuffer::new(),
196 active_buffer: BufferId::A,
197 pending_read: false,
198 })
199 }
200 }
201
202 fn get_inactive_buffer_mut(&mut self) -> &mut [u8] {
203 match self.active_buffer {
204 BufferId::A => self.buffer_b.as_mut_slice(),
205 BufferId::B => self.buffer_a.as_mut_slice(),
206 }
207 }
208
209 fn get_inactive_buffer(&self) -> &[u8] {
210 match self.active_buffer {
211 BufferId::A => self.buffer_b.as_slice(),
212 BufferId::B => self.buffer_a.as_slice(),
213 }
214 }
215
216 fn swap_buffers(&mut self) {
217 self.active_buffer = match self.active_buffer {
218 BufferId::A => BufferId::B,
219 BufferId::B => BufferId::A,
220 };
221 }
222}
223
224#[cfg(target_os = "windows")]
225impl ZeroCopyStreamer for IocpGridManager {
226 fn async_read_chunk(&mut self, offset: u64) -> Result<(), IoError> {
227 if offset % PAGE_SIZE as u64 != 0 {
229 return Err(IoError::MisalignedOffset {
230 offset,
231 required: PAGE_SIZE as u64,
232 });
233 }
234
235 if self.pending_read {
236 return Err(IoError::InvalidState(
237 "Read already in progress. Call poll_completion first.".to_string(),
238 ));
239 }
240
241 self.get_inactive_buffer_mut().fill(0);
243 let _dma_len = self.get_inactive_buffer().len();
244
245 unsafe {
246 let status = crate::directml_bridge::iocp_async_read_ffi(self.handle, offset);
247 if status != crate::directml_bridge::DmlStatus::Success {
248 return Err(IoError::IoError(io::Error::new(
249 io::ErrorKind::Other,
250 status.message(),
251 )));
252 }
253 }
254
255 self.pending_read = true;
256 Ok(())
257 }
258
259 fn poll_completion(&mut self) -> Option<&[u8]> {
260 if !self.pending_read {
261 return None;
262 }
263
264 unsafe {
265 let mut buffer_ptr = std::ptr::null();
266 let mut size = 0usize;
267
268 if crate::directml_bridge::iocp_poll_completion_ffi(
269 self.handle,
270 &mut buffer_ptr,
271 &mut size,
272 ) {
273 self.pending_read = false;
274 self.swap_buffers();
275 Some(self.get_active_buffer())
276 } else {
277 None
278 }
279 }
280 }
281
282 fn get_active_buffer(&self) -> &[u8] {
283 match self.active_buffer {
284 BufferId::A => self.buffer_a.as_slice(),
285 BufferId::B => self.buffer_b.as_slice(),
286 }
287 }
288
289 fn buffer_size(&self) -> usize {
290 DEFAULT_BUFFER_SIZE
291 }
292}
293
294#[cfg(target_os = "windows")]
295impl Drop for IocpGridManager {
296 fn drop(&mut self) {
297 unsafe {
298 crate::directml_bridge::iocp_destroy_ffi(self.handle);
299 }
300 }
301}
302
303#[cfg(target_os = "linux")]
306use libc::{c_void, mlock, O_DIRECT};
307#[cfg(target_os = "linux")]
308use std::os::unix::fs::OpenOptionsExt;
309#[cfg(target_os = "linux")]
310use std::os::unix::io::AsRawFd;
311
312#[cfg(target_os = "linux")]
317pub struct IoUringGridManager {
318 ring: io_uring::IoUring,
319 file: File,
320 buffer_a: DmaBuffer<DEFAULT_BUFFER_SIZE>,
321 buffer_b: DmaBuffer<DEFAULT_BUFFER_SIZE>,
322 active_buffer: BufferId,
323 pending_submission: bool,
324}
325
326#[cfg(target_os = "linux")]
327#[derive(Clone, Copy, PartialEq, Eq)]
328enum BufferId {
329 A,
330 B,
331}
332
333#[cfg(target_os = "linux")]
334impl IoUringGridManager {
335 pub fn new(file_path: &Path) -> Result<Self, IoError> {
340 let file = File::options()
341 .read(true)
342 .custom_flags(O_DIRECT)
343 .open(file_path)
344 .map_err(IoError::FileOpenError)?;
345
346 let ring = io_uring::IoUring::new(8).map_err(IoError::IoError)?;
347
348 let mut manager = Self {
349 ring,
350 file,
351 buffer_a: DmaBuffer::new(),
352 buffer_b: DmaBuffer::new(),
353 active_buffer: BufferId::A,
354 pending_submission: false,
355 };
356
357 manager.pin_buffers()?;
359
360 Ok(manager)
361 }
362
363 fn pin_buffers(&mut self) -> Result<(), IoError> {
365 unsafe {
366 let result_a = mlock(
367 self.buffer_a.as_slice().as_ptr() as *const c_void,
368 self.buffer_a.len(),
369 );
370
371 let result_b = mlock(
372 self.buffer_b.as_slice().as_ptr() as *const c_void,
373 self.buffer_b.len(),
374 );
375
376 if result_a == 0 && result_b == 0 {
377 Ok(())
378 } else {
379 Err(IoError::LockError(
380 "Failed to pin DMA buffers in physical RAM".to_string(),
381 ))
382 }
383 }
384 }
385
386 fn get_inactive_buffer_mut(&mut self) -> &mut [u8] {
387 match self.active_buffer {
388 BufferId::A => self.buffer_b.as_mut_slice(),
389 BufferId::B => self.buffer_a.as_mut_slice(),
390 }
391 }
392
393 fn get_inactive_buffer(&self) -> &[u8] {
394 match self.active_buffer {
395 BufferId::A => self.buffer_b.as_slice(),
396 BufferId::B => self.buffer_a.as_slice(),
397 }
398 }
399
400 fn swap_buffers(&mut self) {
401 self.active_buffer = match self.active_buffer {
402 BufferId::A => BufferId::B,
403 BufferId::B => BufferId::A,
404 };
405 }
406}
407
408#[cfg(target_os = "linux")]
409impl ZeroCopyStreamer for IoUringGridManager {
410 fn async_read_chunk(&mut self, offset: u64) -> Result<(), IoError> {
411 if offset % PAGE_SIZE as u64 != 0 {
413 return Err(IoError::MisalignedOffset {
414 offset,
415 required: PAGE_SIZE as u64,
416 });
417 }
418
419 if self.pending_submission {
420 return Err(IoError::InvalidState(
421 "Read already submitted. Call poll_completion first.".to_string(),
422 ));
423 }
424
425 let read_op = io_uring::opcode::Read::new(
426 io_uring::types::Fd(self.file.as_raw_fd()),
427 self.get_inactive_buffer_mut().as_mut_ptr(),
428 self.get_inactive_buffer().len() as u32,
429 )
430 .offset(offset)
431 .build();
432
433 unsafe {
434 self.ring.submission().push(&read_op).map_err(|e| {
435 IoError::IoError(io::Error::new(io::ErrorKind::Other, e.to_string()))
436 })?;
437 }
438
439 self.pending_submission = true;
440 Ok(())
441 }
442
443 fn poll_completion(&mut self) -> Option<&[u8]> {
444 if !self.pending_submission {
445 return None;
446 }
447
448 match self.ring.submit_and_wait(1) {
449 Ok(_) => {
450 let cqe_result = self.ring.completion().next().map(|cqe| cqe.result());
451 if let Some(res) = cqe_result {
452 self.pending_submission = false;
453
454 if res >= 0 {
455 self.swap_buffers();
456 Some(self.get_active_buffer())
457 } else {
458 None
460 }
461 } else {
462 None
463 }
464 }
465 Err(_) => None,
466 }
467 }
468
469 fn get_active_buffer(&self) -> &[u8] {
470 match self.active_buffer {
471 BufferId::A => self.buffer_a.as_slice(),
472 BufferId::B => self.buffer_b.as_slice(),
473 }
474 }
475
476 fn buffer_size(&self) -> usize {
477 DEFAULT_BUFFER_SIZE
478 }
479}
480
481#[cfg(target_os = "linux")]
482impl Drop for IoUringGridManager {
483 fn drop(&mut self) {
484 unsafe {
485 let _ = libc::munlock(
487 self.buffer_a.as_slice().as_ptr() as *const c_void,
488 self.buffer_a.len(),
489 );
490 let _ = libc::munlock(
491 self.buffer_b.as_slice().as_ptr() as *const c_void,
492 self.buffer_b.len(),
493 );
494 }
495 }
496}
497
498pub struct MmapGridManager {
505 mmap: memmap2::Mmap,
506}
507
508impl MmapGridManager {
509 pub fn new(file_path: &Path) -> Result<Self, IoError> {
513 let file = File::open(file_path).map_err(IoError::FileOpenError)?;
514 let mmap = unsafe { memmap2::Mmap::map(&file) }.map_err(IoError::IoError)?;
515
516 #[cfg(target_os = "linux")]
518 unsafe {
519 libc::madvise(
520 mmap.as_ptr() as *mut libc::c_void,
521 mmap.len(),
522 libc::MADV_SEQUENTIAL | libc::MADV_WILLNEED,
523 );
524 }
525
526 Ok(Self { mmap })
527 }
528
529 pub fn get_slice(&self) -> &[u8] {
531 &self.mmap
532 }
533}
534
535#[cfg(test)]
538mod tests {
539 use super::*;
540 use std::io::Write;
541
542 #[test]
543 fn test_dma_buffer_alignment() {
544 let buffer = DmaBuffer::<4096>::new();
545
546 assert_eq!(buffer.as_slice().as_ptr() as usize % 4096, 0);
548 }
549
550 #[test]
551 #[should_panic(expected = "DMA buffer size must be a multiple of PAGE_SIZE")]
552 fn test_dma_buffer_misaligned_size() {
553 let _buffer = DmaBuffer::<4095>::new();
554 }
555
556 #[test]
557 fn test_dma_buffer_f64_capacity() {
558 let buffer = DmaBuffer::<65536>::new();
559 assert_eq!(buffer.f64_capacity(), 8192); }
561
562 #[cfg(target_os = "windows")]
563 #[test]
564 fn test_iocp_offset_validation() {
565 let temp_dir = std::env::temp_dir();
567 let file_path = temp_dir.join("test_iocp.dat");
568
569 let mut file = File::create(&file_path).unwrap();
570 file.write_all(&vec![0u8; 8192]).unwrap();
572 file.sync_all().unwrap();
573
574 let result = IocpGridManager::new(&file_path);
577 assert!(result.is_err());
578
579 match result {
581 Err(IoError::IoError(e)) => {
582 let error_msg = e.to_string();
583 assert!(
584 error_msg.contains("IOCP implementation stubbed")
585 || error_msg.contains("DirectStorageFailed"),
586 "Expected stub error, got: {}",
587 error_msg
588 );
589 }
590 _ => panic!("Expected IoError::IoError with stub message"),
591 }
592
593 std::fs::remove_file(&file_path).unwrap();
595 }
596
597 #[cfg(target_os = "linux")]
598 #[test]
599 fn test_iouring_offset_validation() {
600 let temp_dir = std::env::temp_dir();
602 let file_path = temp_dir.join("test_iouring.dat");
603
604 let mut file = File::create(&file_path).unwrap();
605 file.write_all(&vec![0u8; 8192]).unwrap();
606 file.sync_all().unwrap();
607
608 let mut manager = IoUringGridManager::new(&file_path).unwrap();
609
610 assert!(manager.async_read_chunk(4096).is_ok());
612
613 let result = manager.async_read_chunk(4095);
615 assert!(matches!(result, Err(IoError::MisalignedOffset { .. })));
616
617 std::fs::remove_file(&file_path).unwrap();
619 }
620}