Skip to main content

ax_net/
raw.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright 2025 KylinSoft Co., Ltd. <https://www.kylinos.cn/>
3// See LICENSES for license details.
4
5//! Raw IP socket implementation for ICMP-style traffic.
6//!
7//! Raw sockets expose packet-oriented access above IP and below TCP/UDP. They
8//! are primarily used by ICMP/ICMPv6 tests and tools, but still share the same
9//! global smoltcp `SocketSet`, route selection, device binding, and readiness
10//! model as UDP/TCP sockets.
11//!
12//! # Packet Format
13//!
14//! smoltcp raw sockets receive complete IP packets. The public raw socket API
15//! returns protocol payloads for normal IPv4/IPv6 raw sockets while preserving
16//! enough packet context for peer filtering and `MSG_PEEK`. Deferred packets
17//! must therefore be stored in a consistent wire-packet form until delivery is
18//! decided.
19//!
20//! # Loopback And Peer Filtering
21//!
22//! Loopback ICMP-style traffic may be delivered through a local fast path. For
23//! connected raw sockets, packets from other peers can be skipped or deferred
24//! without corrupting the smoltcp receive queue format.
25//!
26//! # Locking
27//!
28//! Raw sockets keep their small deferred-packet slots behind IRQ-off spin locks
29//! because packet delivery may be inspected while the protocol executor is
30//! servicing device-originated receive work. These locks are only held around
31//! `Option<Vec<u8>>` swaps and never across route lookup, smoltcp polling, or
32//! userspace buffer I/O.
33
34use alloc::{boxed::Box, sync::Arc, vec};
35use core::{
36    net::{Ipv4Addr, Ipv6Addr, SocketAddr},
37    sync::atomic::{AtomicBool, Ordering},
38    task::Waker,
39};
40
41use ax_io::prelude::*;
42use ax_sync::{SpinLock as Mutex, SpinRwLock as RwLock};
43use axpoll::{ExclusiveRegistrationSink, IoEvents, Pollable, SharedRegistrationSink};
44use axpoll_set::PollSet;
45pub use smoltcp::wire::{IpProtocol, IpVersion};
46use smoltcp::{
47    iface::SocketHandle,
48    socket::raw as smol,
49    storage::PacketMetadata,
50    wire::{Icmpv6Packet, IpAddress, IpListenEndpoint, Ipv4Packet, Ipv4Repr, Ipv6Packet, Ipv6Repr},
51};
52
53use crate::{
54    ConnectStatus, DeferPollWake, NetError, NetResult, RecvFlags, RecvOptions, SOCKET_SET,
55    SendFlags, SendOptions, Shutdown, SocketAddrEx, SocketOps,
56    config::{DeviceBinding, InterfaceId},
57    consts::{RAW_RX_BUF_LEN, RAW_TX_BUF_LEN},
58    general::GeneralOptions,
59    get_control, interface_by_id,
60    ip_tos::apply_ip_tos,
61    options::{Configurable, GetSocketOption, SetSocketOption},
62    request_poll,
63};
64
65enum RawIpHeader {
66    Ipv4(Ipv4Repr),
67    Ipv6(Ipv6Repr),
68}
69
70#[derive(Clone, Copy, PartialEq, Eq)]
71enum RawSocketMode {
72    Raw,
73    IcmpDatagram,
74}
75
76impl RawIpHeader {
77    fn buffer_len(&self) -> usize {
78        match self {
79            Self::Ipv4(header) => header.buffer_len(),
80            Self::Ipv6(header) => header.buffer_len(),
81        }
82    }
83
84    fn emit(&self, buf: &mut [u8]) {
85        match self {
86            Self::Ipv4(header) => header.emit(
87                &mut Ipv4Packet::new_unchecked(buf),
88                &smoltcp::phy::ChecksumCapabilities::ignored(),
89            ),
90            Self::Ipv6(header) => header.emit(&mut Ipv6Packet::new_unchecked(buf)),
91        }
92    }
93}
94
95/// A raw IP socket used for ICMP and ICMPv6 traffic.
96pub struct RawSocket {
97    /// Handle into the global smoltcp socket set.
98    handle: SocketHandle,
99    /// IP version accepted by this socket.
100    ip_version: IpVersion,
101    /// Linux-visible raw or ping-datagram behavior.
102    mode: RawSocketMode,
103    /// Optional local address filter.
104    local_addr: RwLock<Option<IpAddress>>,
105    /// Optional connected peer filter.
106    peer_addr: RwLock<Option<IpAddress>>,
107    /// Locally generated loopback packet waiting to be received.
108    loopback_rx: Mutex<Option<(IpAddress, vec::Vec<u8>)>>,
109    /// Non-peer packet held after filtering without corrupting wire format.
110    deferred_rx: Mutex<Option<(IpAddress, vec::Vec<u8>)>>,
111    /// Optional outgoing TTL/hop-limit override.
112    ttl: RwLock<Option<u8>>,
113    /// Whether recvmsg should report the IPv4 hop limit as ancillary data.
114    recv_ttl: AtomicBool,
115    /// Public read-half closed state.
116    rx_closed: AtomicBool,
117    /// Public write-half closed state.
118    tx_closed: AtomicBool,
119    /// Shared socket options and blocking helpers.
120    general: GeneralOptions,
121    /// Multiplexes protocol and timer wakeups to owned poll registrations.
122    poll_state: Arc<PollSet>,
123}
124
125impl RawSocket {
126    /// Creates a raw socket for the given IP version and protocol.
127    pub fn new(ip_version: IpVersion, ip_protocol: IpProtocol) -> Self {
128        Self::new_with_mode(ip_version, ip_protocol, RawSocketMode::Raw)
129    }
130
131    /// Creates an IPv4 ICMP ping socket with Linux `SOCK_DGRAM` semantics.
132    pub fn new_ipv4_ping() -> Self {
133        Self::new_with_mode(
134            IpVersion::Ipv4,
135            IpProtocol::Icmp,
136            RawSocketMode::IcmpDatagram,
137        )
138    }
139
140    fn new_with_mode(ip_version: IpVersion, ip_protocol: IpProtocol, mode: RawSocketMode) -> Self {
141        let socket_type = match mode {
142            RawSocketMode::Raw => 3,
143            RawSocketMode::IcmpDatagram => 2,
144        };
145        let general = GeneralOptions::new(socket_type, 2, u8::from(ip_protocol) as i32);
146        general.set_device_binding(DeviceBinding::default());
147        Self {
148            handle: SOCKET_SET.add(smol::Socket::new(
149                Some(ip_version),
150                Some(ip_protocol),
151                smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; RAW_RX_BUF_LEN]),
152                smol::PacketBuffer::new(vec![PacketMetadata::EMPTY; 256], vec![0; RAW_TX_BUF_LEN]),
153            )),
154            ip_version,
155            mode,
156            local_addr: RwLock::new(None),
157            peer_addr: RwLock::new(None),
158            loopback_rx: Mutex::new(None),
159            deferred_rx: Mutex::new(None),
160            ttl: RwLock::new(None),
161            recv_ttl: AtomicBool::new(false),
162            rx_closed: AtomicBool::new(false),
163            tx_closed: AtomicBool::new(false),
164            general,
165            poll_state: Arc::new(PollSet::new()),
166        }
167    }
168
169    /// Restricts this socket to one interface for route selection.
170    pub fn bind_device(&self, interface_id: InterfaceId) -> NetResult {
171        if interface_by_id(interface_id).is_none() {
172            return Err(NetError::NoSuchDevice);
173        }
174        self.general.set_device_binding(DeviceBinding {
175            bound_if: Some(interface_id),
176        });
177        Ok(())
178    }
179
180    /// Borrows the underlying smoltcp raw socket by handle.
181    fn with_smol_socket<R>(&self, f: impl FnOnce(&mut smol::Socket) -> R) -> R {
182        SOCKET_SET.with_socket_mut::<smol::Socket, _, _>(self.handle, f)
183    }
184
185    fn outgoing_ip_header(
186        &self,
187        local: IpAddress,
188        remote: IpAddress,
189        next_header: IpProtocol,
190        payload_len: usize,
191        hop_limit: u8,
192    ) -> RawIpHeader {
193        match (self.ip_version, local, remote) {
194            (IpVersion::Ipv4, IpAddress::Ipv4(src_addr), IpAddress::Ipv4(dst_addr)) => {
195                RawIpHeader::Ipv4(Ipv4Repr {
196                    src_addr,
197                    dst_addr,
198                    next_header,
199                    payload_len,
200                    hop_limit,
201                })
202            }
203            (IpVersion::Ipv6, IpAddress::Ipv6(src_addr), IpAddress::Ipv6(dst_addr)) => {
204                RawIpHeader::Ipv6(Ipv6Repr {
205                    src_addr,
206                    dst_addr,
207                    next_header,
208                    payload_len,
209                    hop_limit,
210                })
211            }
212            _ => unreachable!(),
213        }
214    }
215
216    /// Validates that an address belongs to this socket's IP version.
217    fn check_ip_version(&self, addr: IpAddress) -> NetResult<IpAddress> {
218        match (self.ip_version, addr) {
219            (IpVersion::Ipv4, IpAddress::Ipv4(_)) | (IpVersion::Ipv6, IpAddress::Ipv6(_)) => {
220                Ok(addr)
221            }
222            _ => Err(NetError::AddressFamilyUnsupported),
223        }
224    }
225
226    /// Resolves the per-call or connected remote address.
227    fn remote_address(&self, options: &SendOptions) -> NetResult<IpAddress> {
228        match &options.to {
229            Some(addr) => {
230                let remote = addr.clone().into_ip()?;
231                self.check_ip_version(remote.ip().into())
232            }
233            None => (*self.peer_addr.read()).ok_or(NetError::NotConnected),
234        }
235    }
236
237    /// Selects the local source address used for an outgoing raw packet.
238    fn local_address_for(&self, remote: IpAddress) -> NetResult<IpAddress> {
239        if let Some(local) = *self.local_addr.read() {
240            return Ok(local);
241        }
242        if is_loopback_address(remote) {
243            return Ok(remote);
244        }
245        Ok(get_control()
246            .select_route_with_binding(&remote, self.general.device_binding())?
247            .source)
248    }
249
250    /// Splits a complete IP packet into source and bytes returned to userspace.
251    ///
252    /// Linux raw IPv4 receive returns the IP header plus payload, while raw IPv6
253    /// receive returns only the transport payload. The returned slice preserves
254    /// that ABI difference.
255    fn split_packet_for_delivery<'a>(
256        &self,
257        packet: &'a [u8],
258    ) -> NetResult<(IpAddress, &'a [u8], u8)> {
259        match self.ip_version {
260            IpVersion::Ipv4 => {
261                let packet = Ipv4Packet::new_checked(packet).map_err(|_| NetError::InvalidInput)?;
262                let source = IpAddress::Ipv4(packet.src_addr());
263                let hop_limit = packet.hop_limit();
264                let payload = match self.mode {
265                    RawSocketMode::Raw => packet.into_inner(),
266                    RawSocketMode::IcmpDatagram => packet.payload(),
267                };
268                Ok((source, payload, hop_limit))
269            }
270            IpVersion::Ipv6 => {
271                let packet = Ipv6Packet::new_checked(packet).map_err(|_| NetError::InvalidInput)?;
272                Ok((
273                    IpAddress::Ipv6(packet.src_addr()),
274                    packet.payload(),
275                    packet.hop_limit(),
276                ))
277            }
278        }
279    }
280
281    /// Returns whether a received source passes the connected-peer filter.
282    fn source_matches_peer(&self, source: IpAddress) -> bool {
283        self.peer_addr.read().is_none_or(|peer| source == peer)
284    }
285
286    /// Delivers one parsed raw packet to the caller's receive buffer.
287    fn deliver_packet(
288        &self,
289        source: IpAddress,
290        packet: &[u8],
291        hop_limit: u8,
292        dst: &mut (impl Write + IoBufMut),
293        options: &mut RecvOptions<'_>,
294    ) -> NetResult<usize> {
295        if let Some(from) = options.from.as_deref_mut() {
296            *from = SocketAddrEx::Ip(SocketAddr::new(source.into(), 0));
297        }
298        if self.recv_ttl.load(Ordering::Relaxed)
299            && matches!(source, IpAddress::Ipv4(_))
300            && let Some(cmsg) = options.cmsg.as_deref_mut()
301        {
302            cmsg.push(Box::new(crate::IpCmsg::Ipv4Ttl(hop_limit)));
303        }
304
305        let written = dst.write(packet)?;
306        Ok(if options.flags.contains(RecvFlags::TRUNCATE) {
307            packet.len()
308        } else {
309            written
310        })
311    }
312}
313
314fn is_loopback_address(addr: IpAddress) -> bool {
315    match addr {
316        IpAddress::Ipv4(addr) => addr.is_loopback(),
317        IpAddress::Ipv6(addr) => addr.is_loopback(),
318    }
319}
320
321fn icmp_checksum(packet: &[u8]) -> u16 {
322    let mut sum = 0u32;
323    let (chunks, remainder) = packet.as_chunks::<2>();
324    for chunk in chunks {
325        sum += u16::from_be_bytes(*chunk) as u32;
326    }
327    if let Some(&byte) = remainder.first() {
328        sum += u16::from_be_bytes([byte, 0]) as u32;
329    }
330    while sum >> 16 != 0 {
331        sum = (sum & 0xffff) + (sum >> 16);
332    }
333    !(sum as u16)
334}
335
336fn build_loopback_icmp_reply(packet: &[u8]) -> Option<vec::Vec<u8>> {
337    if packet.len() < 8 || packet[0] != 8 || packet[1] != 0 {
338        return None;
339    }
340
341    let mut reply = packet.to_vec();
342    reply[0] = 0;
343    reply[2] = 0;
344    reply[3] = 0;
345    let checksum = icmp_checksum(&reply);
346    reply[2..4].copy_from_slice(&checksum.to_be_bytes());
347    Some(reply)
348}
349
350impl Configurable for RawSocket {
351    fn get_option_inner(&self, option: &mut GetSocketOption) -> NetResult<bool> {
352        use GetSocketOption as O;
353
354        if self.general.get_option_inner(option)? {
355            return Ok(true);
356        }
357
358        match option {
359            O::Ttl(ttl) => {
360                **ttl = (*self.ttl.read()).unwrap_or(64);
361            }
362            O::RecvTtl(enabled) => {
363                **enabled = self.recv_ttl.load(Ordering::Relaxed);
364            }
365            O::SendBuffer(size) => {
366                **size = RAW_TX_BUF_LEN;
367            }
368            O::ReceiveBuffer(size) => {
369                **size = RAW_RX_BUF_LEN;
370            }
371            _ => return Ok(false),
372        }
373        Ok(true)
374    }
375
376    fn set_option_inner(&self, option: SetSocketOption) -> NetResult<bool> {
377        use SetSocketOption as O;
378
379        if self.general.set_option_inner(option)? {
380            return Ok(true);
381        }
382
383        match option {
384            O::Ttl(ttl) => {
385                if *ttl == 0 {
386                    return Err(NetError::InvalidInput);
387                }
388                *self.ttl.write() = Some(*ttl);
389            }
390            O::RecvTtl(enabled) => {
391                self.recv_ttl.store(*enabled, Ordering::Relaxed);
392            }
393            _ => return Ok(false),
394        }
395        Ok(true)
396    }
397}
398
399impl SocketOps for RawSocket {
400    fn bind(&self, local_addr: SocketAddrEx) -> NetResult {
401        let local_addr = local_addr.into_ip()?;
402        let local = self.check_ip_version(local_addr.ip().into())?;
403        *self.local_addr.write() = Some(local);
404        let binding = if local.is_unspecified() {
405            DeviceBinding::default()
406        } else {
407            get_control().local_binding_for(&IpListenEndpoint {
408                addr: Some(local),
409                port: 0,
410            })?
411        };
412        self.general.set_device_binding(binding);
413        Ok(())
414    }
415
416    fn start_connect(&self, remote_addr: SocketAddrEx) -> NetResult<ConnectStatus> {
417        let remote_addr = remote_addr.into_ip()?;
418        let remote = self.check_ip_version(remote_addr.ip().into())?;
419        if self.local_addr.read().is_none() {
420            *self.local_addr.write() = Some(
421                get_control()
422                    .select_route_with_binding(&remote, self.general.device_binding())?
423                    .source,
424            );
425        }
426        *self.peer_addr.write() = Some(remote);
427        let local = (*self.local_addr.read()).expect("raw socket local address");
428        self.general
429            .set_device_binding(get_control().local_binding_for(&IpListenEndpoint {
430                addr: Some(local),
431                port: 0,
432            })?);
433        Ok(ConnectStatus::Connected)
434    }
435
436    fn try_send(&self, mut src: impl Read + IoBuf, options: &mut SendOptions) -> NetResult<usize> {
437        // TODO: MSG_DONTROUTE should bypass the routing table for this datagram.
438        if options.flags.contains(SendFlags::OOB) {
439            return Err(NetError::OperationNotSupported);
440        }
441        if self.tx_closed.load(Ordering::Acquire) {
442            return Err(NetError::BrokenPipe);
443        }
444
445        let remote = self.remote_address(options)?;
446        let local = self.local_address_for(remote)?;
447        let payload_len = src.remaining();
448        let loopback_ipv4 = self.ip_version == IpVersion::Ipv4 && is_loopback_address(remote);
449
450        request_poll();
451        let written = self.with_smol_socket(|socket| {
452            if !socket.can_send() {
453                return Err(NetError::WouldBlock);
454            }
455            let next_header = socket.ip_protocol().expect("raw socket protocol");
456            let hop_limit = (*self.ttl.read()).unwrap_or(64);
457
458            let header =
459                self.outgoing_ip_header(local, remote, next_header, payload_len, hop_limit);
460            let header_len = header.buffer_len();
461
462            let buf = socket
463                .send(header_len + payload_len)
464                .map_err(|_| NetError::WouldBlock)?;
465            header.emit(&mut *buf);
466            let ip_tos = self.general.ip_tos();
467            if ip_tos != 0 {
468                apply_ip_tos(buf, ip_tos);
469            }
470
471            let written = src.read(&mut buf[header_len..])?;
472            if next_header == IpProtocol::Icmpv6 {
473                let (IpAddress::Ipv6(src_addr), IpAddress::Ipv6(dst_addr)) = (local, remote) else {
474                    unreachable!();
475                };
476                Icmpv6Packet::new_unchecked(&mut buf[header_len..])
477                    .fill_checksum(&src_addr, &dst_addr);
478            }
479            if let Some(reply) = loopback_ipv4
480                .then(|| build_loopback_icmp_reply(&buf[header_len..header_len + written]))
481                .flatten()
482            {
483                *self.loopback_rx.lock_irqsave() = Some((local, reply));
484            }
485            Ok(written)
486        })?;
487        request_poll();
488        Ok(written)
489    }
490
491    fn try_recv(
492        &self,
493        mut dst: impl Write + IoBufMut,
494        options: &mut RecvOptions<'_>,
495    ) -> NetResult<usize> {
496        if self.rx_closed.load(Ordering::Acquire) {
497            return Err(NetError::NotConnected);
498        }
499        request_poll();
500        self.with_smol_socket(|socket| {
501            if let Some((source, packet)) = if options.flags.contains(RecvFlags::PEEK) {
502                self.deferred_rx.lock_irqsave().clone()
503            } else {
504                self.deferred_rx.lock_irqsave().take()
505            } {
506                if !self.source_matches_peer(source) {
507                    *self.deferred_rx.lock_irqsave() = Some((source, packet));
508                    return Err(NetError::WouldBlock);
509                }
510                let (_, payload, hop_limit) = self.split_packet_for_delivery(&packet)?;
511                return self.deliver_packet(source, payload, hop_limit, &mut dst, options);
512            }
513
514            if let Some((source, packet)) = if options.flags.contains(RecvFlags::PEEK) {
515                self.loopback_rx.lock_irqsave().clone()
516            } else {
517                self.loopback_rx.lock_irqsave().take()
518            } {
519                if !self.source_matches_peer(source) {
520                    *self.loopback_rx.lock_irqsave() = Some((source, packet));
521                    return Err(NetError::WouldBlock);
522                }
523                return self.deliver_packet(source, &packet, 64, &mut dst, options);
524            }
525
526            let wire_packet = if options.flags.contains(RecvFlags::PEEK) {
527                let packet = socket.peek().map_err(|_| NetError::WouldBlock)?;
528                let (source, ..) = self.split_packet_for_delivery(packet)?;
529                if let Some(peer) = *self.peer_addr.read()
530                    && source != peer
531                {
532                    return Err(NetError::WouldBlock);
533                }
534                packet
535            } else {
536                socket.recv().map_err(|_| NetError::WouldBlock)?
537            };
538            let (source, packet, hop_limit) = self.split_packet_for_delivery(wire_packet)?;
539
540            if !self.source_matches_peer(source) {
541                *self.deferred_rx.lock_irqsave() = Some((source, wire_packet.to_vec()));
542                return Err(NetError::WouldBlock);
543            }
544
545            self.deliver_packet(source, packet, hop_limit, &mut dst, options)
546        })
547    }
548
549    fn local_addr(&self) -> NetResult<SocketAddrEx> {
550        let local = (*self.local_addr.read()).unwrap_or(match self.ip_version {
551            IpVersion::Ipv4 => IpAddress::Ipv4(Ipv4Addr::UNSPECIFIED),
552            IpVersion::Ipv6 => IpAddress::Ipv6(Ipv6Addr::UNSPECIFIED),
553        });
554        Ok(SocketAddrEx::Ip(SocketAddr::new(local.into(), 0)))
555    }
556
557    fn peer_addr(&self) -> NetResult<SocketAddrEx> {
558        let peer = (*self.peer_addr.read()).ok_or(NetError::NotConnected)?;
559        Ok(SocketAddrEx::Ip(SocketAddr::new(peer.into(), 0)))
560    }
561
562    fn shutdown(&self, how: Shutdown) -> NetResult {
563        if how.has_read() {
564            self.rx_closed.store(true, Ordering::Release);
565        }
566        if how.has_write() {
567            self.tx_closed.store(true, Ordering::Release);
568        }
569        Ok(())
570    }
571}
572
573impl Pollable for RawSocket {
574    fn poll(&self) -> IoEvents {
575        request_poll();
576        let mut events = IoEvents::empty();
577        self.with_smol_socket(|socket| {
578            events.set(
579                IoEvents::IN,
580                !self.rx_closed.load(Ordering::Acquire) && socket.can_recv(),
581            );
582            events.set(
583                IoEvents::OUT,
584                !self.tx_closed.load(Ordering::Acquire) && socket.can_send(),
585            );
586        });
587        events.set(
588            IoEvents::IN,
589            events.contains(IoEvents::IN)
590                || self
591                    .loopback_rx
592                    .lock_irqsave()
593                    .as_ref()
594                    .is_some_and(|(source, _)| self.source_matches_peer(*source))
595                || self
596                    .deferred_rx
597                    .lock_irqsave()
598                    .as_ref()
599                    .is_some_and(|(source, _)| self.source_matches_peer(*source)),
600        );
601        events
602    }
603
604    unsafe fn register_shared(&self, sink: &mut dyn SharedRegistrationSink, events: IoEvents) {
605        unsafe { sink.register_shared(&self.poll_state, events) };
606        self.arm_poll_sources(events);
607    }
608
609    unsafe fn register_exclusive(
610        &self,
611        sink: &mut dyn ExclusiveRegistrationSink,
612        events: IoEvents,
613    ) {
614        unsafe { sink.register_exclusive(&self.poll_state, events) };
615        self.arm_poll_sources(events);
616    }
617}
618
619impl RawSocket {
620    fn arm_poll_sources(&self, events: IoEvents) {
621        self.with_smol_socket(|socket| {
622            if events.contains(IoEvents::IN) {
623                socket.register_recv_waker(&Waker::from(Arc::new(DeferPollWake {
624                    poll: self.poll_state.clone(),
625                    ready: IoEvents::IN,
626                })));
627            }
628            if events.contains(IoEvents::OUT) {
629                socket.register_send_waker(&Waker::from(Arc::new(DeferPollWake {
630                    poll: self.poll_state.clone(),
631                    ready: IoEvents::OUT,
632                })));
633            }
634        });
635        if events.intersects(IoEvents::IN | IoEvents::OUT) {
636            self.general
637                .register_waker(&Waker::from(Arc::new(DeferPollWake {
638                    poll: self.poll_state.clone(),
639                    ready: events,
640                })));
641        }
642    }
643}
644
645impl Drop for RawSocket {
646    fn drop(&mut self) {
647        self.shutdown(Shutdown::Both).ok();
648        SOCKET_SET.remove(self.handle);
649    }
650}