1use 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
95pub struct RawSocket {
97 handle: SocketHandle,
99 ip_version: IpVersion,
101 mode: RawSocketMode,
103 local_addr: RwLock<Option<IpAddress>>,
105 peer_addr: RwLock<Option<IpAddress>>,
107 loopback_rx: Mutex<Option<(IpAddress, vec::Vec<u8>)>>,
109 deferred_rx: Mutex<Option<(IpAddress, vec::Vec<u8>)>>,
111 ttl: RwLock<Option<u8>>,
113 recv_ttl: AtomicBool,
115 rx_closed: AtomicBool,
117 tx_closed: AtomicBool,
119 general: GeneralOptions,
121 poll_state: Arc<PollSet>,
123}
124
125impl RawSocket {
126 pub fn new(ip_version: IpVersion, ip_protocol: IpProtocol) -> Self {
128 Self::new_with_mode(ip_version, ip_protocol, RawSocketMode::Raw)
129 }
130
131 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 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 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 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 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 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 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 fn source_matches_peer(&self, source: IpAddress) -> bool {
283 self.peer_addr.read().is_none_or(|peer| source == peer)
284 }
285
286 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 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}