1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
use libc::{msync, MS_SYNC};
use std::io::{self, Result};
use std::ptr;
use libc::{close, munmap, PROT_READ, PROT_WRITE};
use crate::shmem_sys;
pub(crate) struct ShmemWriter {
ptr: *mut u8,
current_ptr: *mut u8,
size: usize,
fd: i32,
name: String,
}
// SAFETY: the only non-auto fields are `ptr`/`current_ptr`, raw pointers into an
// mmap'd shared-memory region whose address is fixed for the handle's lifetime.
// Moving or sharing the handle across threads is sound; data-race freedom on the
// mapped bytes is the caller's responsibility (as for any shared memory) — the
// higher-level wrappers serialize writes where needed.
unsafe impl Send for ShmemWriter {}
unsafe impl Sync for ShmemWriter {}
impl ShmemWriter {
/// Opens and maps a shared memory region for read/write access.
pub fn new(name: &str, size: usize, unlock_mapped_memory: bool) -> Result<Self> {
// Open existing shared memory (read/write)
let fd = shmem_sys::open(name, libc::O_RDWR)?;
// Map the memory region for read/write
let ptr = shmem_sys::map(fd, size, PROT_READ | PROT_WRITE, !unlock_mapped_memory, name)?;
let ptr_u8 = ptr as *mut u8;
Ok(Self { ptr: ptr_u8, current_ptr: ptr_u8, size, fd, name: name.to_string() })
}
unsafe fn unmap(&mut self) {
if munmap(self.ptr as *mut _, self.size) != 0 {
tracing::error!("munmap failed: {:?}", io::Error::last_os_error());
} else {
self.ptr = ptr::null_mut();
self.size = 0;
tracing::trace!("Unmapped shared memory '{}'", self.name);
}
}
/// Writes data to the shared memory, starting at the specified offset
///
/// # Type Parameters
/// * `T` - The element type of the slice (e.g., u8, u64)
///
/// # Arguments
/// * `offset` - Byte offset from the start of shared memory where data should be written
/// * `data` - A slice of data to write to shared memory
///
/// # Returns
/// * `Ok(())` - If data was successfully written
/// * `Err` - If data size exceeds shared memory capacity or msync fails
pub fn write_at<T>(&self, offset: usize, data: &[T]) -> Result<()> {
let byte_size = std::mem::size_of_val(data);
if byte_size > self.size {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"Data size ({} bytes) exceeds shared memory capacity ({}) for '{}'",
byte_size, self.size, self.name
),
));
}
unsafe {
ptr::copy_nonoverlapping(data.as_ptr() as *const u8, self.ptr.add(offset), byte_size);
// Force changes to be flushed to the shared memory
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
if msync(self.ptr as *mut _, self.size, MS_SYNC /*| MS_INVALIDATE*/) != 0 {
return Err(io::Error::last_os_error());
}
}
Ok(())
}
/// Writes data to the shared memory, always from the start
///
/// # Type Parameters
/// * `T` - The element type of the slice (e.g., u8, u64)
///
/// # Arguments
/// * `data` - A slice of data to write to shared memory
///
/// # Returns
/// * `Ok(())` - If data was successfully written
/// * `Err` - If data size exceeds shared memory capacity or msync fails
pub fn append_input<T>(&mut self, data: &[T]) -> Result<()> {
let byte_size = std::mem::size_of_val(data);
if byte_size > self.size {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"Data size ({} bytes) exceeds shared memory capacity ({}) for '{}'",
byte_size, self.size, self.name
),
));
}
unsafe {
ptr::copy_nonoverlapping(data.as_ptr() as *const u8, self.current_ptr, byte_size);
// Force changes to be flushed to the shared memory
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
if msync(self.ptr as *mut _, self.size, MS_SYNC) != 0 {
return Err(io::Error::last_os_error());
}
self.current_ptr = self.current_ptr.add(byte_size);
}
Ok(())
}
/// Writes data to the shared memory as a ring buffer, handling wraparound automatically
///
/// Uses internal pointer tracking with automatic wraparound.
///
/// # Type Parameters
/// * `T` - The element type of the slice (e.g., u8, u64)
///
/// # Arguments
/// * `data` - A slice of data to write to shared memory
#[inline]
pub fn write_ring_buffer<T>(&mut self, data: &[T]) -> Result<()> {
let byte_size = std::mem::size_of_val(data);
let data_ptr = data.as_ptr() as *const u8;
unsafe {
let current_offset = self.current_ptr.offset_from(self.ptr) as usize;
// Check if data wraps around the buffer
if current_offset + byte_size > self.size {
// Split write: first part to end of buffer, second part from start
let first_part_size = self.size - current_offset;
let second_part_size = byte_size - first_part_size;
// Write first part to end of buffer
ptr::copy_nonoverlapping(data_ptr, self.current_ptr, first_part_size);
// Write second part to start of buffer
ptr::copy_nonoverlapping(data_ptr.add(first_part_size), self.ptr, second_part_size);
// Update current_ptr to point after the second part
self.current_ptr = self.ptr.add(second_part_size);
} else {
// Write contiguously
ptr::copy_nonoverlapping(data_ptr, self.current_ptr, byte_size);
// Update current_ptr, wrapping if at end
self.current_ptr = self.current_ptr.add(byte_size);
let new_offset = self.current_ptr.offset_from(self.ptr) as usize;
if new_offset == self.size {
self.current_ptr = self.ptr;
}
}
// Force changes to be flushed to the shared memory
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
if msync(self.ptr as *mut _, self.size, MS_SYNC) != 0 {
return Err(io::Error::last_os_error());
}
}
Ok(())
}
/// Reads a u64 from shared memory at a specific offset (in bytes)
///
/// # Arguments
/// * `offset` - Byte offset from the start of shared memory (must be 8-byte aligned)
///
/// # Safety
/// This method assumes that:
/// - The shared memory contains at least `offset + 8` bytes of valid data
/// - The offset should be aligned to 8 bytes
///
/// # Returns
/// * The u64 value read from the specified offset (in native endianness)
#[inline]
pub fn read_u64_at(&self, offset: usize) -> u64 {
debug_assert_eq!(offset % 8, 0, "Offset must be 8-byte aligned");
unsafe { (self.ptr.add(offset) as *const u64).read() }
}
/// Writes a u64 to shared memory at a specific offset (in bytes)
///
/// # Arguments
/// * `offset` - Byte offset from the start of shared memory (must be 8-byte aligned)
/// * `value` - The u64 value to write
///
/// # Safety
/// This method assumes that:
/// - The shared memory contains at least `offset + 8` bytes of valid data
/// - The offset is 8-byte aligned for optimal performance
///
/// # Returns
/// * `Ok(())` - If the value was written and flushed
/// * `Err` - If the `msync` flushing the write fails (Linux only)
#[inline]
pub fn write_u64_at(&self, offset: usize, value: u64) -> Result<()> {
debug_assert_eq!(offset % 8, 0, "Offset must be 8-byte aligned");
unsafe {
(self.ptr.add(offset) as *mut u64).write(value);
// Force changes to be flushed to the shared memory
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
if msync(self.ptr as *mut _, self.size, MS_SYNC) != 0 {
return Err(io::Error::last_os_error());
}
}
Ok(())
}
/// Resets the internal pointer used for appending to the start of the buffer.
pub fn reset(&mut self) {
self.current_ptr = self.ptr;
}
}
impl Drop for ShmemWriter {
fn drop(&mut self) {
unsafe {
self.unmap();
close(self.fd);
}
}
}
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
#[cfg(test)]
mod tests {
use super::*;
use crate::ShmemReader;
use std::ffi::CString;
/// Unique per-test segment name (cargo runs tests as threads in one process).
fn seg_name(tag: &str) -> String {
format!("ZISK_unittest_{}_{tag}", std::process::id())
}
/// Create a fresh, zero-filled `/dev/shm` segment of `size` bytes.
fn create_segment(name: &str, size: usize) {
let c = CString::new(name).unwrap();
unsafe {
libc::shm_unlink(c.as_ptr()); // drop any stale leftover
let fd = libc::shm_open(c.as_ptr(), libc::O_CREAT | libc::O_RDWR, 0o600);
assert!(fd >= 0, "shm_open(create) failed for {name}");
assert_eq!(libc::ftruncate(fd, size as libc::off_t), 0, "ftruncate failed");
libc::close(fd);
}
}
fn unlink_segment(name: &str) {
let c = CString::new(name).unwrap();
unsafe { libc::shm_unlink(c.as_ptr()) };
}
#[test]
fn writer_reader_u64_round_trip() {
let name = seg_name("u64rt");
create_segment(&name, 4096);
{
let w = ShmemWriter::new(&name, 4096, true).unwrap();
w.write_u64_at(0, 0xDEAD_BEEF).unwrap();
w.write_u64_at(8, 42).unwrap();
}
let r = ShmemReader::new(&name, 4096).unwrap();
assert_eq!(r.read_u64_at(0), 0xDEAD_BEEF);
assert_eq!(r.read_u64_at(8), 42);
unlink_segment(&name);
}
#[test]
fn write_at_offset_is_visible_to_reader() {
let name = seg_name("writeat");
create_segment(&name, 4096);
let w = ShmemWriter::new(&name, 4096, true).unwrap();
w.write_at(16, &[1u64, 2, 3]).unwrap();
let r = ShmemReader::new(&name, 4096).unwrap();
assert_eq!([r.read_u64_at(16), r.read_u64_at(24), r.read_u64_at(32)], [1, 2, 3]);
unlink_segment(&name);
}
#[test]
fn write_at_rejects_payload_larger_than_segment() {
let name = seg_name("cap");
create_segment(&name, 64);
let w = ShmemWriter::new(&name, 64, true).unwrap();
assert!(w.write_at(0, &[0u8; 128]).is_err());
unlink_segment(&name);
}
#[test]
fn ring_buffer_wraps_around_the_end() {
let name = seg_name("ring");
create_segment(&name, 32); // 4 u64 slots
let mut w = ShmemWriter::new(&name, 32, true).unwrap();
w.write_ring_buffer(&[1u64, 2, 3]).unwrap(); // slots 0,8,16; cursor at 24
w.write_ring_buffer(&[4u64, 5]).unwrap(); // 4 at 24, wraps, 5 at 0
let r = ShmemReader::new(&name, 32).unwrap();
assert_eq!(r.read_u64_at(24), 4, "last slot before wrap");
assert_eq!(r.read_u64_at(0), 5, "wrapped to start");
assert_eq!(r.read_u64_at(8), 2, "untouched from first write");
unlink_segment(&name);
}
#[test]
fn new_fails_for_missing_segment() {
let name = seg_name("missing");
unlink_segment(&name); // ensure it does not exist
assert!(ShmemWriter::new(&name, 4096, true).is_err());
}
}