1use 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>>, 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 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 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 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 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 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 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 let mut free_bufs: Vec<u16> = (0..self.qd as u16).rev().collect();
126
127 while completed < requests.len() {
128 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 self.ring
144 .submit_and_wait(1)
145 .context("io_uring submit_and_wait")?;
146 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 stream_sync(stream)?;
177 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 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 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}