1use alloc::{boxed::Box, vec::Vec};
22use core::{
23 any::Any,
24 fmt::{self, Debug},
25 net::SocketAddr,
26 time::Duration,
27};
28
29use ax_io::prelude::*;
30use axpoll::{ExclusiveRegistrationSink, IoEvents, Pollable, SharedRegistrationSink};
31use bitflags::bitflags;
32use enum_dispatch::enum_dispatch;
33
34#[cfg(feature = "vsock")]
35use crate::vsock::{VsockAddr, VsockSocket};
36use crate::{
37 NetError, NetResult,
38 options::{Configurable, GetSocketOption, SetSocketOption, UnixCredentials},
39 raw::RawSocket,
40 tcp::TcpSocket,
41 udp::UdpSocket,
42 unix::{UnixSocket, UnixSocketAddr},
43};
44
45#[derive(Clone, Debug)]
47pub enum SocketAddrEx {
48 Ip(SocketAddr),
50 Unix(UnixSocketAddr),
52 #[cfg(feature = "vsock")]
54 Vsock(VsockAddr),
55}
56
57impl SocketAddrEx {
58 pub fn into_ip(self) -> NetResult<SocketAddr> {
60 match self {
61 SocketAddrEx::Ip(addr) => Ok(addr),
62 SocketAddrEx::Unix(_) => Err(NetError::AddressFamilyUnsupported),
63 #[cfg(feature = "vsock")]
64 SocketAddrEx::Vsock(_) => Err(NetError::AddressFamilyUnsupported),
65 }
66 }
67
68 pub fn into_unix(self) -> NetResult<UnixSocketAddr> {
70 match self {
71 SocketAddrEx::Unix(addr) => Ok(addr),
72 SocketAddrEx::Ip(_) => Err(NetError::AddressFamilyUnsupported),
73 #[cfg(feature = "vsock")]
74 SocketAddrEx::Vsock(_) => Err(NetError::AddressFamilyUnsupported),
75 }
76 }
77
78 #[cfg(feature = "vsock")]
80 pub fn into_vsock(self) -> NetResult<VsockAddr> {
81 match self {
82 SocketAddrEx::Ip(_) => Err(NetError::AddressFamilyUnsupported),
83 SocketAddrEx::Unix(_) => Err(NetError::AddressFamilyUnsupported),
84 SocketAddrEx::Vsock(addr) => Ok(addr),
85 }
86 }
87}
88
89bitflags! {
90 #[derive(Default, Debug, Clone, Copy)]
97 pub struct SendFlags: u32 {
98 const OOB = 0x01;
100 const DONTROUTE = 0x04;
103 const DONTWAIT = 0x40;
106 const EOR = 0x80;
108 const CONFIRM = 0x800;
110 const NOSIGNAL = 0x4000;
113 const MORE = 0x8000;
115 }
116}
117
118bitflags! {
119 #[derive(Default, Debug, Clone, Copy)]
123 pub struct RecvFlags: u32 {
124 const PEEK = 0x01;
126 const TRUNCATE = 0x02;
130 const OOB = 0x04;
133 const DONTWAIT = 0x40;
136 }
137}
138
139pub trait CMsgPayload: Any + Send + Sync {
145 fn clone_box(&self) -> Box<dyn CMsgPayload>;
147 fn into_any(self: Box<Self>) -> Box<dyn Any + Send + Sync>;
149}
150impl<T: Any + Send + Sync + Clone> CMsgPayload for T {
151 fn clone_box(&self) -> Box<dyn CMsgPayload> {
152 Box::new(self.clone())
153 }
154 fn into_any(self: Box<Self>) -> Box<dyn Any + Send + Sync> {
155 self
156 }
157}
158impl core::fmt::Debug for dyn CMsgPayload {
161 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
162 f.write_str("CMsgPayload { .. }")
163 }
164}
165
166pub type CMsgData = Box<dyn CMsgPayload>;
168
169#[derive(Debug, Clone, Copy, PartialEq, Eq)]
171pub enum IpCmsg {
172 Ipv4Ttl(u8),
174 Ipv4Tos(u8),
176 Ipv6TrafficClass(u8),
178}
179
180#[derive(Debug, Clone, PartialEq, Eq)]
183pub enum SocketCmsg {
184 Credentials(UnixCredentials),
186 Timestamp(Duration),
188}
189
190#[derive(Default, Debug)]
194pub struct SendOptions {
195 pub to: Option<SocketAddrEx>,
197 pub flags: SendFlags,
199 pub cmsg: Vec<CMsgData>,
201 pub sender_credentials: Option<UnixCredentials>,
203}
204
205#[derive(Default)]
209pub struct RecvOptions<'a> {
210 pub from: Option<&'a mut SocketAddrEx>,
212 pub flags: RecvFlags,
214 pub cmsg: Option<&'a mut Vec<CMsgData>>,
216 pub truncated: Option<&'a mut bool>,
218}
219impl Debug for RecvOptions<'_> {
220 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
221 f.debug_struct("RecvOptions")
222 .field("from", &self.from)
223 .field("flags", &self.flags)
224 .finish()
225 }
226}
227
228#[derive(Debug, Clone, Copy)]
230pub enum Shutdown {
231 Read,
233 Write,
235 Both,
237}
238
239#[derive(Clone, Copy, Debug, Eq, PartialEq)]
241pub enum ConnectStatus {
242 Connected,
244 InProgress,
246}
247
248#[derive(Clone, Copy, Debug, Eq, PartialEq)]
250pub struct SocketWaitPolicy {
251 pub nonblocking: bool,
253 pub timeout: Option<Duration>,
255}
256
257fn socket_wait_policy(
258 socket: &(impl Configurable + ?Sized),
259 send: bool,
260 extra_nonblocking: bool,
261) -> NetResult<SocketWaitPolicy> {
262 let mut nonblocking = false;
263 socket.get_option(GetSocketOption::NonBlocking(&mut nonblocking))?;
264 let mut timeout = Duration::ZERO;
265 if send {
266 socket.get_option(GetSocketOption::SendTimeout(&mut timeout))?;
267 } else {
268 socket.get_option(GetSocketOption::ReceiveTimeout(&mut timeout))?;
269 }
270 Ok(SocketWaitPolicy {
271 nonblocking: nonblocking || extra_nonblocking,
272 timeout: (!timeout.is_zero()).then_some(timeout),
273 })
274}
275impl Shutdown {
276 pub fn has_read(&self) -> bool {
278 matches!(self, Shutdown::Read | Shutdown::Both)
279 }
280
281 pub fn has_write(&self) -> bool {
283 matches!(self, Shutdown::Write | Shutdown::Both)
284 }
285}
286
287#[enum_dispatch]
289pub trait SocketOps: Configurable {
290 fn bind(&self, local_addr: SocketAddrEx) -> NetResult;
292 fn start_connect(&self, remote_addr: SocketAddrEx) -> NetResult<ConnectStatus>;
294 fn connect_status(&self) -> NetResult<ConnectStatus> {
296 Ok(ConnectStatus::Connected)
297 }
298
299 fn listen(&self, _backlog: usize) -> NetResult {
301 Err(NetError::OperationNotSupported)
302 }
303 fn is_listening(&self) -> bool {
305 false
306 }
307 fn try_accept(&self) -> NetResult<Socket> {
309 Err(NetError::OperationNotSupported)
310 }
311
312 fn try_send(&self, src: impl Read + IoBuf, options: &mut SendOptions) -> NetResult<usize>;
314 fn try_recv(
316 &self,
317 dst: impl Write + IoBufMut,
318 options: &mut RecvOptions<'_>,
319 ) -> NetResult<usize>;
320 fn recv_available(&self) -> NetResult<usize> {
322 Err(NetError::OperationNotSupported)
323 }
324
325 fn local_addr(&self) -> NetResult<SocketAddrEx>;
327 fn peer_addr(&self) -> NetResult<SocketAddrEx>;
329
330 fn shutdown(&self, how: Shutdown) -> NetResult;
332
333 fn send_wait_policy(&self, extra_nonblocking: bool) -> NetResult<SocketWaitPolicy> {
335 socket_wait_policy(self, true, extra_nonblocking)
336 }
337
338 fn recv_wait_policy(&self, extra_nonblocking: bool) -> NetResult<SocketWaitPolicy> {
340 socket_wait_policy(self, false, extra_nonblocking)
341 }
342}
343
344impl<T: SocketOps + ?Sized> SocketOps for Box<T> {
345 fn bind(&self, local_addr: SocketAddrEx) -> NetResult {
346 (**self).bind(local_addr)
347 }
348
349 fn start_connect(&self, remote_addr: SocketAddrEx) -> NetResult<ConnectStatus> {
350 (**self).start_connect(remote_addr)
351 }
352
353 fn connect_status(&self) -> NetResult<ConnectStatus> {
354 (**self).connect_status()
355 }
356
357 fn listen(&self, backlog: usize) -> NetResult {
358 (**self).listen(backlog)
359 }
360
361 fn is_listening(&self) -> bool {
362 (**self).is_listening()
363 }
364
365 fn try_accept(&self) -> NetResult<Socket> {
366 (**self).try_accept()
367 }
368
369 fn try_send(&self, src: impl Read + IoBuf, options: &mut SendOptions) -> NetResult<usize> {
370 (**self).try_send(src, options)
371 }
372
373 fn try_recv(
374 &self,
375 dst: impl Write + IoBufMut,
376 options: &mut RecvOptions<'_>,
377 ) -> NetResult<usize> {
378 (**self).try_recv(dst, options)
379 }
380
381 fn recv_available(&self) -> NetResult<usize> {
382 (**self).recv_available()
383 }
384
385 fn local_addr(&self) -> NetResult<SocketAddrEx> {
386 (**self).local_addr()
387 }
388
389 fn peer_addr(&self) -> NetResult<SocketAddrEx> {
390 (**self).peer_addr()
391 }
392
393 fn shutdown(&self, how: Shutdown) -> NetResult {
394 (**self).shutdown(how)
395 }
396}
397
398#[enum_dispatch(Configurable, SocketOps)]
400pub enum Socket {
401 Udp(Box<UdpSocket>),
403 Tcp(Box<TcpSocket>),
405 Raw(Box<RawSocket>),
407 Unix(Box<UnixSocket>),
409 #[cfg(feature = "vsock")]
411 Vsock(Box<VsockSocket>),
412}
413
414impl From<UdpSocket> for Socket {
415 fn from(socket: UdpSocket) -> Self {
416 Self::Udp(Box::new(socket))
417 }
418}
419
420impl From<TcpSocket> for Socket {
421 fn from(socket: TcpSocket) -> Self {
422 Self::Tcp(Box::new(socket))
423 }
424}
425
426impl From<UnixSocket> for Socket {
427 fn from(socket: UnixSocket) -> Self {
428 Self::Unix(Box::new(socket))
429 }
430}
431
432#[cfg(feature = "vsock")]
433impl From<VsockSocket> for Socket {
434 fn from(socket: VsockSocket) -> Self {
435 Self::Vsock(Box::new(socket))
436 }
437}
438
439impl Pollable for Socket {
440 fn poll(&self) -> IoEvents {
441 match self {
442 Socket::Tcp(tcp) => tcp.poll(),
443 Socket::Udp(udp) => udp.poll(),
444 Socket::Raw(raw) => raw.poll(),
445 Socket::Unix(unix) => unix.poll(),
446 #[cfg(feature = "vsock")]
447 Socket::Vsock(vsock) => vsock.poll(),
448 }
449 }
450
451 unsafe fn register_shared(&self, sink: &mut dyn SharedRegistrationSink, events: IoEvents) {
452 match self {
453 Socket::Tcp(tcp) => unsafe { tcp.register_shared(sink, events) },
454 Socket::Udp(udp) => unsafe { udp.register_shared(sink, events) },
455 Socket::Raw(raw) => unsafe { raw.register_shared(sink, events) },
456 Socket::Unix(unix) => unsafe { unix.register_shared(sink, events) },
457 #[cfg(feature = "vsock")]
458 Socket::Vsock(vsock) => unsafe { vsock.register_shared(sink, events) },
459 }
460 }
461
462 unsafe fn register_exclusive(
463 &self,
464 sink: &mut dyn ExclusiveRegistrationSink,
465 events: IoEvents,
466 ) {
467 match self {
468 Socket::Tcp(tcp) => unsafe { tcp.register_exclusive(sink, events) },
469 Socket::Udp(udp) => unsafe { udp.register_exclusive(sink, events) },
470 Socket::Raw(raw) => unsafe { raw.register_exclusive(sink, events) },
471 Socket::Unix(unix) => unsafe { unix.register_exclusive(sink, events) },
472 #[cfg(feature = "vsock")]
473 Socket::Vsock(vsock) => unsafe { vsock.register_exclusive(sink, events) },
474 }
475 }
476}