Skip to content

Commit 920c49c

Browse files
committed
Add a test of sending a memory file descriptor over a socket.
1 parent 23f0db8 commit 920c49c

3 files changed

Lines changed: 176 additions & 0 deletions

File tree

Cargo.lock

Lines changed: 14 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

wgpu-remote/Cargo.toml

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -25,5 +25,13 @@ optional = true
2525
version = "1.0"
2626
features = ["derive"]
2727

28+
[target.'cfg(unix)'.dev-dependencies]
29+
nix = { version = "0.29", features = ["fs", "mman", "process", "socket", "uio"] }
30+
31+
[[test]]
32+
name = "unix-send-shmem"
33+
path = "tests/unix-send-shmem.rs"
34+
harness = false
35+
2836
[lints]
2937
workspace = true
Lines changed: 154 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,154 @@
1+
/// Exercise sending shared memory via Unix domain sockets.
2+
///
3+
/// This isn't really a test of the `wgpu-remote` crate's
4+
/// functionality; it's a verification that the system actually
5+
/// behaves as `wgpu-remote` expects.
6+
///
7+
/// If it turns out that all Unix systems behave identically, then
8+
/// this is a waste of time. But if it turns out that there are bugs,
9+
/// quirks, or limits, then this serves to detect and document them.
10+
11+
use std::os::fd::{self, AsRawFd, FromRawFd};
12+
13+
use nix::sys::{memfd, mman, socket, wait};
14+
use nix::unistd;
15+
16+
fn main() {
17+
let (sender_socket, receiver_socket) = socket::socketpair(
18+
socket::AddressFamily::Unix,
19+
socket::SockType::Stream,
20+
None,
21+
socket::SockFlag::empty(),
22+
).unwrap();
23+
24+
// Start the sender.
25+
// Safety: we are single-threaded.
26+
let fork = unsafe { unistd::fork() }.unwrap();
27+
let unistd::ForkResult::Parent { child: sender_pid } = fork else {
28+
drop(receiver_socket);
29+
sender(sender_socket);
30+
// never returns
31+
};
32+
33+
// Start the receiver.
34+
// Safety: we are single-threaded.
35+
let fork = unsafe { unistd::fork() }.unwrap();
36+
let unistd::ForkResult::Parent { child: receiver_pid } = fork else {
37+
drop(sender_socket);
38+
receiver(receiver_socket);
39+
// never returns
40+
};
41+
42+
drop(sender_socket);
43+
drop(receiver_socket);
44+
45+
let status = wait::waitpid(sender_pid, None).unwrap();
46+
assert_eq!(status, wait::WaitStatus::Exited(sender_pid, 42));
47+
48+
let status = wait::waitpid(receiver_pid, None).unwrap();
49+
assert_eq!(status, wait::WaitStatus::Exited(receiver_pid, 43));
50+
}
51+
52+
const LEN: usize = 10;
53+
54+
fn sender(socket: fd::OwnedFd) -> ! {
55+
// Create a file descriptor referring to zeroed memory.
56+
//
57+
// The `nix` source code suggests `memfd_create` is available on
58+
// Linux, Android, and FreeBSD.
59+
let mem_fd = memfd::memfd_create(
60+
c"unix-send-shmem",
61+
memfd::MemFdCreateFlag::empty(),
62+
).unwrap();
63+
64+
// Set the file's size.
65+
unistd::ftruncate(&mem_fd, LEN as _).unwrap();
66+
67+
// Map the memory file descriptor into our address space.
68+
//
69+
// Safety: We're not requesting an address, our offset is aligned,
70+
// and our flags are reasonable.
71+
let mapped = unsafe {
72+
mman::mmap(
73+
None,
74+
std::num::NonZeroUsize::new(LEN).unwrap(),
75+
mman::ProtFlags::PROT_READ | mman::ProtFlags::PROT_WRITE,
76+
mman::MapFlags::MAP_SHARED,
77+
&mem_fd,
78+
0,
79+
).unwrap()
80+
};
81+
82+
// Write a message to the memory.
83+
//
84+
// Safety: the memory is there and zeroed, and `u8` has no
85+
// requirement alignments.
86+
let mapped = unsafe { mapped.cast::<[u8; LEN]>().as_mut() };
87+
mapped.copy_from_slice(b"Greetings!");
88+
89+
// Send the file descriptor to the other process.
90+
socket::sendmsg::<()>(
91+
socket.as_raw_fd(),
92+
&[std::io::IoSlice::new(b"Hi")],
93+
&[socket::ControlMessage::ScmRights(&[mem_fd.as_raw_fd()])],
94+
socket::MsgFlags::empty(),
95+
None, // address
96+
).unwrap();
97+
98+
std::process::exit(42);
99+
}
100+
101+
fn receiver(socket: fd::OwnedFd) -> ! {
102+
// Receive the message from the sender.
103+
let mut data_buf = [0_u8; 10];
104+
let mut slices = [std::io::IoSliceMut::new(&mut data_buf)];
105+
let mut fd_buf = nix::cmsg_space!([std::os::fd::RawFd; 20]);
106+
let msg = socket::recvmsg::<()>(
107+
socket.as_raw_fd(),
108+
&mut slices, // iov (ordinary bytes)
109+
Some(&mut fd_buf), // cmsg_buffer
110+
socket::MsgFlags::empty(),
111+
).unwrap();
112+
113+
// Get the memory file descriptor out of the message.
114+
assert_eq!(msg.bytes, 2);
115+
let mut cmsgs = msg.cmsgs().unwrap();
116+
let mut fds;
117+
match cmsgs.next() {
118+
Some(socket::ControlMessageOwned::ScmRights(f)) => { fds = f; }
119+
Some(other) => panic!("Got unexpected control message: {other:#?}"),
120+
None => panic!("didn't get any control messages"),
121+
}
122+
assert_eq!(cmsgs.next(), None);
123+
let raw_mem_fd = fds.pop().unwrap();
124+
// Safety: we just got this fd from a control message, so it should be open.
125+
let mem_fd = unsafe { fd::OwnedFd::from_raw_fd(raw_mem_fd) };
126+
assert!(fds.is_empty());
127+
128+
// Check the bytes that were transmitted.
129+
assert_eq!(&slices[0][..2], b"Hi");
130+
131+
// Map the memory file descriptor into our address space.
132+
//
133+
// Safety: We're not requesting an address, our offset is aligned,
134+
// and our flags are reasonable.
135+
let mapped = unsafe {
136+
mman::mmap(
137+
None,
138+
std::num::NonZeroUsize::new(LEN).unwrap(),
139+
mman::ProtFlags::PROT_READ,
140+
mman::MapFlags::MAP_SHARED,
141+
&mem_fd,
142+
0,
143+
).unwrap()
144+
};
145+
146+
// Check the message in the memory.
147+
//
148+
// Safety: the memory is there and initialized, and `u8` has no
149+
// requirement alignments.
150+
let mapped = unsafe { mapped.cast::<[u8; LEN]>().as_ref() };
151+
assert_eq!(mapped, b"Greetings!");
152+
153+
std::process::exit(43);
154+
}

0 commit comments

Comments
 (0)