spark_storage/backend/
io_uring.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2//
3// Phase-3 production storage backend: `io_uring` (IORING_SETUP_SQPOLL +
4// IORING_REGISTER_BUFFERS) + per-buffer `CudaEvent` for safe reuse across
5// async H→D DMAs. Per-buffer events let us keep QD≥8 in flight without the
6// per-op `cuStreamSynchronize` that throttled the POSIX backend to QD=1.
7
8use anyhow::{Context, Result, bail};
9use io_uring::{IoUring, opcode, types};
10use std::ffi::c_void;
11use std::os::fd::RawFd;
12
13use super::{ReadRequest, StorageBackend};
14use crate::cuda_min::{CudaEvent, PinnedBuffer, copy_h_to_d_async, stream_sync};
15use crate::group::{GroupKey, GroupLayout};
16use crate::layout::Layout;
17
18pub struct IoUringBackend {
19    layout: Layout,
20    ring: IoUring,
21    buffers: Vec<PinnedBuffer>,
22    events: Vec<Option<CudaEvent>>, // event per buffer, None = idle
23    qd: usize,
24}
25
26impl IoUringBackend {
27    pub fn new(layout: Layout, qd: usize) -> Result<Self> {
28        if qd == 0 {
29            bail!("queue depth must be ≥ 1");
30        }
31        // SQPOLL: kernel polls SQ; idle 2s before parking.
32        let ring = IoUring::builder()
33            .setup_sqpoll(2_000)
34            .build(qd as u32)
35            .context("io_uring build")?;
36
37        let group_bytes = layout.group_bytes() as usize;
38        let mut buffers = Vec::with_capacity(qd);
39        for _ in 0..qd {
40            buffers.push(PinnedBuffer::new(group_bytes)?);
41        }
42        // Register the pinned host buffers with io_uring for zero-copy
43        // direct-IO. After this, ReadFixed at index `i` lands in `buffers[i]`.
44        let iovecs: Vec<libc::iovec> = buffers
45            .iter()
46            .map(|b| libc::iovec {
47                iov_base: b.ptr,
48                iov_len: b.bytes,
49            })
50            .collect();
51        unsafe {
52            ring.submitter()
53                .register_buffers(&iovecs)
54                .context("register_buffers")?;
55        }
56        let events: Vec<Option<CudaEvent>> = (0..qd).map(|_| None).collect();
57        Ok(Self {
58            layout,
59            ring,
60            buffers,
61            events,
62            qd,
63        })
64    }
65
66    pub fn layout(&self) -> &Layout {
67        &self.layout
68    }
69
70    /// Test helper: drop the page cache for the layer files so subsequent
71    /// reads actually hit NVMe.
72    pub fn drop_pagecache(&self) {
73        for layer in 0..self.layout.spec.num_layers {
74            let fd = self.layout.fd(layer);
75            unsafe { libc::posix_fadvise(fd, 0, 0, libc::POSIX_FADV_DONTNEED) };
76        }
77    }
78
79    /// Wait for the previous DMA out of `buf_idx` to complete (if any) so
80    /// we can reuse the buffer for a new io_uring read.
81    fn wait_buffer_free(&mut self, buf_idx: usize) -> Result<()> {
82        if let Some(ev) = self.events[buf_idx].take() {
83            ev.sync()?;
84        }
85        Ok(())
86    }
87
88    /// Submit one read request into `buf_idx` and return its user_data tag.
89    fn submit_read(
90        &mut self,
91        fd: RawFd,
92        offset: u64,
93        bytes: u32,
94        buf_idx: u16,
95        user_data: u64,
96    ) -> Result<()> {
97        let buf_ptr = self.buffers[buf_idx as usize].ptr as *mut u8;
98        let read_e = opcode::ReadFixed::new(types::Fd(fd), buf_ptr, bytes, buf_idx)
99            .offset(offset)
100            .build()
101            .user_data(user_data);
102        unsafe {
103            self.ring
104                .submission()
105                .push(&read_e)
106                .map_err(|_| anyhow::anyhow!("io_uring SQ full"))?;
107        }
108        Ok(())
109    }
110}
111
112impl StorageBackend for IoUringBackend {
113    fn read(&mut self, requests: &[ReadRequest], stream: u64) -> Result<()> {
114        let bytes = self.layout.group_bytes() as u32;
115        // user_data layout: high 16 bits = req index, low 16 bits = buf index.
116        // (We never submit > 65k requests in one batch.)
117        if requests.len() > u16::MAX as usize {
118            bail!("io_uring batch too large: {}", requests.len());
119        }
120
121        let mut next_submit = 0;
122        let mut completed = 0;
123        // Buffer ownership: free buffers form a stack; busy ones are claimed
124        // by an in-flight read until its CQE arrives.
125        let mut free_bufs: Vec<u16> = (0..self.qd as u16).rev().collect();
126
127        while completed < requests.len() {
128            // Submit while we have a free buffer and pending requests.
129            while next_submit < requests.len() {
130                let Some(&buf_idx) = free_bufs.last() else {
131                    break;
132                };
133                self.wait_buffer_free(buf_idx as usize)?;
134                free_bufs.pop();
135                let req = &requests[next_submit];
136                let fd = self.layout.fd(req.group.layer);
137                let off = self.layout.offset(req.group);
138                let user = ((next_submit as u64) << 16) | (buf_idx as u64);
139                self.submit_read(fd, off, bytes, buf_idx, user)?;
140                next_submit += 1;
141            }
142            // Submit and wait for at least one completion.
143            self.ring
144                .submit_and_wait(1)
145                .context("io_uring submit_and_wait")?;
146            // Drain everything that's ready.
147            let cq = self.ring.completion();
148            for cqe in cq {
149                let user = cqe.user_data();
150                let buf_idx = (user & 0xFFFF) as u16;
151                let req_idx = (user >> 16) as usize;
152                let result = cqe.result();
153                if result < 0 {
154                    bail!("io_uring read failed for req {req_idx}: errno {}", -result);
155                }
156                if result as u32 != bytes {
157                    bail!("io_uring short read: req {req_idx} got {result}, expected {bytes}");
158                }
159                let req = &requests[req_idx];
160                let buf = &self.buffers[buf_idx as usize];
161                copy_h_to_d_async(
162                    req.dst_dev_ptr,
163                    buf.ptr as *const c_void,
164                    bytes as usize,
165                    stream,
166                )?;
167                let ev = CudaEvent::new()?;
168                ev.record(stream)?;
169                self.events[buf_idx as usize] = Some(ev);
170                free_bufs.push(buf_idx);
171                completed += 1;
172            }
173        }
174        // After all reads have produced device data, finalise the stream
175        // (matches PosixBackend semantics: at return, the stream is synced).
176        stream_sync(stream)?;
177        // Drop now-completed events; they are useful only across calls.
178        for slot in self.events.iter_mut() {
179            *slot = None;
180        }
181        Ok(())
182    }
183
184    fn write_from_host(&mut self, key: GroupKey, src: &[u8]) -> Result<()> {
185        let bytes = self.layout.group_bytes() as usize;
186        if src.len() != bytes {
187            bail!(
188                "write_from_host: src len {} != group bytes {bytes}",
189                src.len()
190            );
191        }
192        // Stage through buffer 0 — pinned + page-aligned for O_DIRECT.
193        self.wait_buffer_free(0)?;
194        unsafe {
195            std::ptr::copy_nonoverlapping(src.as_ptr(), self.buffers[0].ptr as *mut u8, bytes);
196        }
197        let fd = self.layout.fd(key.layer);
198        let off = self.layout.offset(key) as i64;
199        let n = unsafe { libc::pwrite(fd, self.buffers[0].ptr, bytes, off) };
200        if n != bytes as isize {
201            bail!(
202                "pwrite {bytes}@{off} returned {n}, errno {}",
203                std::io::Error::last_os_error()
204            );
205        }
206        Ok(())
207    }
208
209    fn group_layout(&self) -> GroupLayout {
210        self.layout.spec
211    }
212}
213
214#[cfg(test)]
215mod tests {
216    use super::*;
217    use crate::cuda_min::{CudaCtx, DeviceBuffer, copy_d_to_h_async};
218    use crate::group::{GroupKey, GroupLayout, KvKind};
219    use std::path::PathBuf;
220
221    fn tempdir(name: &str) -> PathBuf {
222        let p = std::env::temp_dir().join(format!("atlas-iouring-{}-{}", name, std::process::id()));
223        let _ = std::fs::remove_dir_all(&p);
224        std::fs::create_dir_all(&p).unwrap();
225        p
226    }
227
228    #[test]
229    #[ignore = "requires GPU"]
230    fn write_then_read_round_trip() {
231        let _ctx = CudaCtx::new(0).expect("cuda init");
232        let dir = tempdir("rt");
233        let spec = GroupLayout::new(1, 4, 1, 16, 128, 2, 4096);
234        let layout = Layout::create(&dir, spec).unwrap();
235        let mut backend = IoUringBackend::new(layout, 4).unwrap();
236        let bytes = backend.layout().group_bytes() as usize;
237        // Three different patterns at three different keys to exercise SQ depth.
238        let patterns: Vec<(GroupKey, Vec<u8>)> = (0..4u32)
239            .map(|b| {
240                let k = GroupKey::new(0, b, 0, KvKind::K);
241                let pat: Vec<u8> = (0..bytes)
242                    .map(|i| ((i + b as usize) & 0xFF) as u8)
243                    .collect();
244                (k, pat)
245            })
246            .collect();
247        for (k, p) in &patterns {
248            backend.write_from_host(*k, p).unwrap();
249        }
250        backend.drop_pagecache();
251        let dev: Vec<DeviceBuffer> = patterns
252            .iter()
253            .map(|_| DeviceBuffer::new(bytes).unwrap())
254            .collect();
255        let reqs: Vec<ReadRequest> = patterns
256            .iter()
257            .zip(&dev)
258            .map(|((k, _), d)| ReadRequest {
259                group: *k,
260                dst_dev_ptr: d.ptr,
261            })
262            .collect();
263        backend.read(&reqs, _ctx.stream).unwrap();
264        for ((_, expected), d) in patterns.iter().zip(&dev) {
265            let mut got = vec![0_u8; bytes];
266            copy_d_to_h_async(got.as_mut_ptr() as *mut c_void, d.ptr, bytes, _ctx.stream).unwrap();
267            stream_sync(_ctx.stream).unwrap();
268            assert_eq!(&got, expected);
269        }
270        std::fs::remove_dir_all(&dir).ok();
271    }
272}