Skip to main content

compio_driver/sys/op/socket/
unix.rs

1use std::{net::Shutdown, num::NonZeroU32};
2
3use compio_buf::{IoBufExt, IoBufMutExt};
4use rustix::{
5    io::close,
6    net::{
7        AddressFamily, Protocol, RecvAncillaryBuffer, SendAncillaryBuffer, SocketAddrAny,
8        SocketType, acceptfrom_with, bind, connect, listen, recv, recvfrom, recvmsg, send, sendmsg,
9        sendmsg_addr, sendto, shutdown, socket_with,
10    },
11};
12
13use crate::{PollFirst, sys::op::*};
14
15impl<S: AsFd> Accept<S> {
16    pub(crate) fn call(&mut self) -> io::Result<usize> {
17        let (owned, addr) = acceptfrom_with(self.fd.as_fd(), SOCKET_FLAG)?;
18        let fd = owned.as_raw_fd();
19        let socket: Socket2 = owned.into();
20
21        if cfg!(apple) {
22            socket.set_cloexec(true)?;
23            socket.set_nonblocking(true)?;
24        }
25
26        copy_addr_from(&mut self.buffer, &mut self.addr_len, addr);
27        self.accepted_fd = Some(socket);
28
29        Ok(fd as usize)
30    }
31}
32
33impl<S: AsFd> Connect<S> {
34    pub(crate) fn call(&self) -> io::Result<usize> {
35        connect(&self.fd, &SockAddrArg(&self.addr))?;
36        Ok(0)
37    }
38}
39
40impl<T: IoBuf, S: AsFd> Send<T, S> {
41    pub(crate) fn call(&mut self) -> io::Result<usize> {
42        send(self.fd.as_fd(), self.buffer.as_init(), self.flags).map_err(Into::into)
43    }
44}
45
46impl<T: IoBuf, S: AsFd> SendTo<T, S> {
47    pub(crate) fn call(&self) -> io::Result<usize> {
48        sendto(
49            self.header.fd.as_fd(),
50            self.buffer.as_init(),
51            self.header.flags,
52            &SockAddrArg(&self.header.addr),
53        )
54        .map_err(Into::into)
55    }
56}
57
58impl<T: IoVectoredBuf, S: AsFd> SendVectored<T, S> {
59    pub(crate) fn call(&self, control: &mut SendVectoredControl) -> io::Result<usize> {
60        let mut anc = SendAncillaryBuffer::default();
61
62        sendmsg(
63            self.fd.as_fd(),
64            io_slice(&control.slices),
65            &mut anc,
66            self.flags,
67        )
68        .map_err(Into::into)
69    }
70}
71
72impl<T: IoVectoredBuf, S: AsFd> SendToVectored<T, S> {
73    pub(crate) fn call(&mut self, control: &mut SendMsgControl) -> io::Result<usize> {
74        let addr = SockAddrArg(&self.header.addr);
75        let mut anc = SendAncillaryBuffer::default();
76        let buf = io_slice(&control.slices);
77
78        sendmsg_addr(
79            self.header.fd.as_fd(),
80            &addr,
81            buf,
82            &mut anc,
83            self.header.flags,
84        )
85        .map_err(Into::into)
86    }
87}
88
89impl<T: IoVectoredBuf, C: IoBuf, S: AsFd> SendMsg<T, C, S> {
90    pub(crate) fn call(&mut self, control: &mut SendMsgControl) -> io::Result<usize> {
91        // Both rustix and nix expose api that uses structured AncillaryBuffer
92        // building, no way to just throw in an ancillary buf. Fallback to libc here.
93        syscall!(libc::sendmsg(
94            self.fd.as_fd().as_raw_fd(),
95            &control.msg,
96            self.flags.bits() as _,
97        ))
98    }
99}
100
101impl<T: IoBufMut, S: AsFd> Recv<T, S> {
102    pub(crate) fn call(&mut self) -> io::Result<usize> {
103        let (_, len) = recv(self.fd.as_fd(), self.buffer.as_uninit(), self.flags)?;
104
105        Ok(len)
106    }
107}
108
109impl<T: IoVectoredBufMut, S: AsFd> RecvVectored<T, S> {
110    pub(crate) fn call(&mut self, control: &mut RecvVectoredControl) -> io::Result<usize> {
111        let res = recvmsg(
112            self.fd.as_fd(),
113            io_slice_mut(&mut control.slices),
114            &mut RecvAncillaryBuffer::default(),
115            self.flags,
116        )?;
117
118        // Kernel may truncate and return a larger-than-buffer size
119        Ok(res.bytes.min(self.buffer.total_capacity()))
120    }
121}
122
123impl<S: AsFd> RecvFromHeader<S> {
124    pub fn set_addr(&mut self, addr: Option<SocketAddrAny>) {
125        copy_addr_from(&mut self.addr, &mut self.addr_len, addr)
126    }
127}
128
129impl<T: IoBufMut, S: AsFd> RecvFrom<T, S> {
130    pub(crate) fn call(&mut self) -> io::Result<usize> {
131        let (_, len, addr) = recvfrom(&self.header.fd, self.buffer.as_uninit(), self.header.flags)?;
132
133        self.header.set_addr(addr);
134
135        Ok(len.min(self.buffer.buf_capacity()))
136    }
137}
138
139impl<T: IoVectoredBufMut, C: IoBufMut, S: AsFd> RecvMsg<T, C, S> {
140    pub(crate) fn call(&mut self, control: &mut RecvMsgControl) -> io::Result<usize> {
141        let res = syscall!(libc::recvmsg(
142            self.header.fd.as_fd().as_raw_fd(),
143            &raw mut control.msg,
144            self.header.flags.bits() as _,
145        ))?;
146
147        self.update_control(control);
148
149        Ok(res)
150    }
151}
152
153impl<T: IoVectoredBufMut, S: AsFd> RecvFromVectored<T, S> {
154    pub(crate) fn call(&mut self, control: &mut RecvMsgControl) -> io::Result<usize> {
155        let res = recvmsg(
156            &self.header.fd,
157            io_slice_mut(&mut control.slices),
158            &mut RecvAncillaryBuffer::default(),
159            self.header.flags,
160        )?;
161
162        self.header.set_addr(res.address);
163
164        Ok(res.bytes)
165    }
166}
167
168/// Create a socket.
169pub struct CreateSocket {
170    pub(crate) domain: AddressFamily,
171    pub(crate) socket_type: SocketType,
172    pub(crate) protocol: Option<Protocol>,
173    pub(crate) opened_fd: Option<Socket2>,
174}
175
176impl CreateSocket {
177    /// Create [`CreateSocket`].
178    pub fn new(domain: i32, socket_type: i32, protocol: i32) -> Self {
179        let domain = AddressFamily::from_raw(domain as _);
180        let socket_type = SocketType::from_raw(socket_type as _);
181        let protocol = NonZeroU32::new(protocol as _).map(Protocol::from_raw);
182
183        Self {
184            domain,
185            socket_type,
186            protocol,
187            opened_fd: None,
188        }
189    }
190
191    pub(crate) fn call(&mut self) -> io::Result<usize> {
192        let owned = socket_with(self.domain, self.socket_type, SOCKET_FLAG, self.protocol)?;
193        let fd = owned.as_raw_fd();
194        let socket: Socket2 = owned.into();
195
196        #[cfg(apple)]
197        {
198            socket.set_cloexec(true)?;
199            socket.set_nosigpipe(true)?;
200            socket.set_nonblocking(true)?;
201        }
202
203        self.opened_fd = Some(socket);
204        Ok(fd as _)
205    }
206}
207
208impl IntoInner for CreateSocket {
209    type Inner = Socket2;
210
211    fn into_inner(self) -> Self::Inner {
212        self.opened_fd.expect("socket not created")
213    }
214}
215
216/// Bind a socket to an address.
217pub struct Bind<S> {
218    pub(crate) fd: S,
219    pub(crate) addr: SockAddr,
220}
221
222impl<S> Bind<S> {
223    /// Create [`Bind`].
224    pub fn new(fd: S, addr: SockAddr) -> Self {
225        Self { fd, addr }
226    }
227}
228
229impl<S: AsFd> Bind<S> {
230    pub(crate) fn call(&self) -> io::Result<usize> {
231        bind(self.fd.as_fd(), &SockAddrArg(&self.addr))?;
232        Ok(0)
233    }
234}
235
236/// Listen for connections on a socket.
237pub struct Listen<S> {
238    pub(crate) fd: S,
239    pub(crate) backlog: i32,
240}
241
242impl<S> Listen<S> {
243    /// Create [`Listen`].
244    pub fn new(fd: S, backlog: i32) -> Self {
245        Self { fd, backlog }
246    }
247}
248
249impl<S: AsFd> Listen<S> {
250    pub(crate) fn call(&self) -> io::Result<usize> {
251        listen(self.fd.as_fd(), self.backlog)?;
252        Ok(0)
253    }
254}
255
256/// Shutdown a socket.
257pub struct ShutdownSocket<S> {
258    pub(crate) fd: S,
259    pub(crate) how: Shutdown,
260}
261
262impl<S> ShutdownSocket<S> {
263    /// Create [`ShutdownSocket`].
264    pub fn new(fd: S, how: Shutdown) -> Self {
265        Self { fd, how }
266    }
267}
268
269impl<S: AsFd> ShutdownSocket<S> {
270    #[cfg(io_uring)]
271    pub(crate) fn how(&self) -> i32 {
272        match self.how {
273            Shutdown::Write => libc::SHUT_WR,
274            Shutdown::Read => libc::SHUT_RD,
275            Shutdown::Both => libc::SHUT_RDWR,
276        }
277    }
278
279    pub(crate) fn call(&mut self) -> io::Result<usize> {
280        let how = match self.how {
281            Shutdown::Write => rustix::net::Shutdown::Write,
282            Shutdown::Read => rustix::net::Shutdown::Read,
283            Shutdown::Both => rustix::net::Shutdown::Both,
284        };
285        shutdown(&self.fd, how)?;
286        Ok(0)
287    }
288}
289
290impl CloseSocket {
291    pub(crate) fn call(&mut self) -> io::Result<usize> {
292        unsafe { close(self.fd.as_raw_fd()) };
293        Ok(0)
294    }
295}
296
297/// Accept a connection.
298pub struct Accept<S> {
299    pub(crate) fd: S,
300    pub(crate) buffer: SockAddrStorage,
301    pub(crate) addr_len: socklen_t,
302    pub(crate) accepted_fd: Option<Socket2>,
303    pub(crate) poll_first: bool,
304}
305
306impl<S> Accept<S> {
307    /// Create [`Accept`].
308    pub fn new(fd: S) -> Self {
309        let buffer = SockAddrStorage::zeroed();
310        let addr_len = buffer.size_of();
311        Self {
312            fd,
313            buffer,
314            addr_len,
315            accepted_fd: None,
316            poll_first: false,
317        }
318    }
319}
320
321impl<S> PollFirst for Accept<S> {
322    fn poll_first(&mut self) {
323        self.poll_first = true;
324    }
325}
326
327impl<S> IntoInner for Accept<S> {
328    type Inner = (Socket2, SockAddr);
329
330    fn into_inner(mut self) -> Self::Inner {
331        let socket = self.accepted_fd.take().expect("socket not accepted");
332        (socket, unsafe { SockAddr::new(self.buffer, self.addr_len) })
333    }
334}
335
336#[doc(hidden)]
337pub struct RecvVectoredControl {
338    pub(crate) msg: libc::msghdr,
339    #[allow(dead_code)]
340    pub(crate) slices: Vec<SysSlice>,
341}
342
343impl Default for RecvVectoredControl {
344    fn default() -> Self {
345        Self {
346            msg: unsafe { std::mem::zeroed() },
347            slices: Vec::new(),
348        }
349    }
350}
351
352impl<T: IoVectoredBufMut, S> RecvVectored<T, S> {
353    pub(crate) fn init_control(&mut self, ctrl: &mut RecvVectoredControl) {
354        ctrl.slices = self.buffer.sys_slices_mut();
355        ctrl.msg.msg_iov = ctrl.slices.as_mut_ptr() as _;
356        ctrl.msg.msg_iovlen = ctrl.slices.len() as _;
357    }
358}
359
360#[doc(hidden)]
361pub struct SendVectoredControl {
362    pub(crate) msg: libc::msghdr,
363    #[allow(dead_code)]
364    pub(crate) slices: Vec<SysSlice>,
365}
366
367impl Default for SendVectoredControl {
368    fn default() -> Self {
369        Self {
370            msg: unsafe { std::mem::zeroed() },
371            slices: Vec::new(),
372        }
373    }
374}
375
376impl<T: IoVectoredBuf, S> SendVectored<T, S> {
377    pub(crate) fn init_control(&mut self, ctrl: &mut SendVectoredControl) {
378        ctrl.slices = self.buffer.sys_slices();
379        ctrl.msg.msg_iov = ctrl.slices.as_ptr() as _;
380        ctrl.msg.msg_iovlen = ctrl.slices.len() as _;
381    }
382}
383
384#[doc(hidden)]
385pub struct SendMsgControl {
386    pub(crate) msg: libc::msghdr,
387    #[allow(dead_code)]
388    pub(crate) slices: Multi<SysSlice>,
389}
390
391impl<S: AsFd> SendToHeader<S> {
392    #[allow(dead_code)]
393    pub(crate) fn create_control(
394        &mut self,
395        ctrl: &mut SendMsgControl,
396        slices: impl Into<Multi<SysSlice>>,
397    ) {
398        ctrl.msg.msg_name = self.addr.as_ptr() as _;
399        ctrl.msg.msg_namelen = self.addr.len();
400        ctrl.slices = slices.into();
401        ctrl.msg.msg_iov = ctrl.slices.as_mut_ptr() as _;
402        ctrl.msg.msg_iovlen = ctrl.slices.len() as _;
403    }
404}
405
406impl Default for SendMsgControl {
407    fn default() -> Self {
408        Self {
409            msg: unsafe { std::mem::zeroed() },
410            slices: Multi::new(),
411        }
412    }
413}
414
415impl<T: IoVectoredBuf, C: IoBuf, S> SendMsg<T, C, S> {
416    pub(crate) fn init_control(&mut self, ctrl: &mut SendMsgControl) {
417        ctrl.slices = self.buffer.sys_slices().into();
418        match self.addr.as_ref() {
419            Some(addr) => {
420                ctrl.msg.msg_name = addr.as_ptr() as _;
421                ctrl.msg.msg_namelen = addr.len();
422            }
423            None => {
424                ctrl.msg.msg_name = std::ptr::null_mut();
425                ctrl.msg.msg_namelen = 0;
426            }
427        }
428        ctrl.msg.msg_iov = ctrl.slices.as_ptr() as _;
429        ctrl.msg.msg_iovlen = ctrl.slices.len() as _;
430        ctrl.msg.msg_control = self.control.buf_ptr() as _;
431        ctrl.msg.msg_controllen = self.control.buf_len() as _;
432    }
433}
434
435#[doc(hidden)]
436pub struct RecvMsgControl {
437    pub(crate) msg: libc::msghdr,
438    #[allow(dead_code)]
439    pub(crate) slices: Multi<SysSlice>,
440}
441
442impl Default for RecvMsgControl {
443    fn default() -> Self {
444        Self {
445            msg: unsafe { std::mem::zeroed() },
446            slices: Multi::new(),
447        }
448    }
449}
450
451impl<T: IoVectoredBufMut, C: IoBufMut, S> RecvMsg<T, C, S> {
452    pub(crate) fn init_control(&mut self, ctrl: &mut RecvMsgControl) {
453        ctrl.slices = Multi::from_vec(self.buffer.sys_slices_mut());
454        ctrl.msg.msg_name = &raw mut self.header.addr as _;
455        ctrl.msg.msg_namelen = self.header.addr.size_of() as _;
456        ctrl.msg.msg_iov = ctrl.slices.as_mut_ptr() as _;
457        ctrl.msg.msg_iovlen = ctrl.slices.len() as _;
458        ctrl.msg.msg_control = self.control.buf_mut_ptr() as _;
459        ctrl.msg.msg_controllen = self.control.buf_capacity() as _;
460    }
461
462    pub(crate) fn update_control(&mut self, control: &RecvMsgControl) {
463        self.header.addr_len = control.msg.msg_namelen as _;
464        self.control_len = control.msg.msg_controllen as _;
465        self.return_flags = ReturnFlags::from_bits_retain(control.msg.msg_flags as _);
466    }
467}