Skip to main content

starnix_core/vfs/socket/
socket_netlink.rs

1// Copyright 2022 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use crate::security::{self, AuditLogger, AuditMessage, AuditRequest};
6use crate::vfs::socket::{SockOptValue, SocketDomain};
7use futures::channel::mpsc::{
8    UnboundedReceiver, UnboundedSender, {self},
9};
10use linux_uapi::{AUDIT_GET, NETLINK_GET_STRICT_CHK, audit_status};
11use netlink::messaging::{
12    AccessControl, MessageWithPermission, NetlinkContext, NetlinkMessageWithCreds, Permission,
13    Sender, UnparsedNetlinkMessage,
14};
15use netlink::multicast_groups::{
16    InvalidLegacyGroupsError, InvalidModernGroupError, LegacyGroups, ModernGroup,
17    NoMappingFromModernToLegacyGroupError, SingleLegacyGroup,
18};
19use netlink::protocol_family::NetlinkClient;
20use netlink::protocol_family::route::NetlinkRouteClient;
21use netlink::protocol_family::sock_diag::NetlinkSockDiagClient;
22use netlink::{NETLINK_LOG_TAG, NewClientError};
23use netlink_packet_core::{
24    ErrorMessage, NETLINK_HEADER_LEN, NLMSG_ERROR, NetlinkBuffer, NetlinkDeserializable,
25    NetlinkHeader, NetlinkMessage, NetlinkPayload, NetlinkSerializable,
26};
27use netlink_packet_generic::message::EmptyDeserializeOptions as EmptyDeserializeGenlOptions;
28use netlink_packet_route::{RouteNetlinkMessage, RouteNetlinkMessageParseMode};
29use netlink_packet_sock_diag::SockDiagRequest;
30use netlink_packet_sock_diag::message::EmptyDeserializeOptions as EmptyDeserializeSockDiagOptions;
31use netlink_packet_utils::{DecodeError, Emitable as _};
32use starnix_sync::{
33    AuditNetlinkClientAuditResponseLock, LockDepGuard, LockDepMutex, NetlinkSocketInnerLock,
34    UEventNetlinkSocketDeviceListenerKeyLock,
35};
36use std::io::Write;
37use std::marker::PhantomData;
38use std::num::{NonZeroI32, NonZeroU32};
39use std::sync::Arc;
40use zerocopy::{FromBytes, IntoBytes};
41
42use crate::device::kobject::{Device, UEventAction, UEventContext, flatten_uevent_properties};
43use crate::device::{DeviceListener, DeviceListenerKey};
44use crate::task::{CurrentTask, EventHandler, Kernel, WaitCanceler, WaitQueue, Waiter};
45use crate::vfs::buffers::{
46    AncillaryData, InputBuffer, Message, MessageQueue, MessageReadInfo, OutputBuffer,
47    UnixControlData, VecInputBuffer,
48};
49use crate::vfs::socket::{
50    GenericMessage, GenericNetlinkClientHandle, Socket, SocketAddress, SocketHandle,
51    SocketMessageFlags, SocketOps, SocketPeer, SocketShutdownFlags, SocketType,
52};
53use starnix_logging::{log_debug, log_error, log_warn, track_stub};
54use starnix_uapi::auth::{CAP_AUDIT_CONTROL, CAP_AUDIT_WRITE, CAP_NET_ADMIN, Credentials};
55use starnix_uapi::errors::Errno;
56use starnix_uapi::vfs::FdEvents;
57use starnix_uapi::{
58    AF_NETLINK, NETLINK_ADD_MEMBERSHIP, NETLINK_AUDIT, NETLINK_CONNECTOR, NETLINK_CRYPTO,
59    NETLINK_DNRTMSG, NETLINK_DROP_MEMBERSHIP, NETLINK_ECRYPTFS, NETLINK_FIB_LOOKUP,
60    NETLINK_FIREWALL, NETLINK_GENERIC, NETLINK_IP6_FW, NETLINK_ISCSI, NETLINK_KOBJECT_UEVENT,
61    NETLINK_NETFILTER, NETLINK_NFLOG, NETLINK_RDMA, NETLINK_ROUTE, NETLINK_SCSITRANSPORT,
62    NETLINK_SELINUX, NETLINK_SMC, NETLINK_SOCK_DIAG, NETLINK_USERSOCK, NETLINK_XFRM, NLM_F_MULTI,
63    NLMSG_DONE, SO_PASSCRED, SO_PROTOCOL, SO_RCVBUF, SO_RCVBUFFORCE, SO_SNDBUF, SO_SNDBUFFORCE,
64    SO_TIMESTAMP, SOL_SOCKET, errno, error, nlmsghdr, sockaddr_nl, socklen_t, ucred,
65};
66
67// From netlink/socket.go in gVisor.
68pub const SOCKET_MIN_SIZE: usize = 4 << 10;
69pub const SOCKET_DEFAULT_SIZE: usize = 16 * 1024;
70pub const SOCKET_MAX_SIZE: usize = 4 << 20;
71
72// From linux/socket.go in gVisor.
73const SOL_NETLINK: u32 = 270;
74
75pub fn new_netlink_socket(
76    kernel: &Arc<Kernel>,
77    socket_type: SocketType,
78    family: NetlinkFamily,
79) -> Result<Box<dyn SocketOps>, Errno> {
80    log_debug!(tag = NETLINK_LOG_TAG; "Creating {:?} Netlink Socket", family);
81    if socket_type != SocketType::Datagram && socket_type != SocketType::Raw {
82        return error!(ESOCKTNOSUPPORT);
83    }
84
85    let ops: Box<dyn SocketOps> = match family {
86        NetlinkFamily::KobjectUevent => Box::new(UEventNetlinkSocket::default()),
87        NetlinkFamily::Route => Box::new(new_route_socket(kernel)?),
88        NetlinkFamily::Generic => Box::new(GenericNetlinkSocket::new(kernel)?),
89        NetlinkFamily::SockDiag => Box::new(new_sock_diag_socket(kernel)?),
90        NetlinkFamily::Audit => Box::new(AuditNetlinkSocket::new(kernel)?),
91        NetlinkFamily::Usersock
92        | NetlinkFamily::Firewall
93        | NetlinkFamily::Nflog
94        | NetlinkFamily::Xfrm
95        | NetlinkFamily::Selinux
96        | NetlinkFamily::Iscsi
97        | NetlinkFamily::FibLookup
98        | NetlinkFamily::Connector
99        | NetlinkFamily::Netfilter
100        | NetlinkFamily::Ip6Fw
101        | NetlinkFamily::Dnrtmsg
102        | NetlinkFamily::Scsitransport
103        | NetlinkFamily::Ecryptfs
104        | NetlinkFamily::Rdma
105        | NetlinkFamily::Crypto
106        | NetlinkFamily::Smc => Box::new(StubbedNetlinkSocket::new(family)),
107        NetlinkFamily::Invalid => return error!(EINVAL),
108    };
109    Ok(ops)
110}
111
112#[derive(Default, Debug, Clone, PartialEq, Eq)]
113#[repr(C)]
114pub struct NetlinkAddress {
115    pid: u32,
116    groups: u32,
117}
118
119impl NetlinkAddress {
120    pub fn new(pid: u32, groups: u32) -> Self {
121        NetlinkAddress { pid, groups }
122    }
123
124    pub fn set_pid_if_zero(&mut self, pid: i32) {
125        if self.pid == 0 {
126            self.pid = pid as u32;
127        }
128    }
129
130    pub fn to_bytes(&self) -> Vec<u8> {
131        sockaddr_nl { nl_family: AF_NETLINK, nl_pid: self.pid, nl_pad: 0, nl_groups: self.groups }
132            .as_bytes()
133            .to_vec()
134    }
135}
136
137#[derive(Debug, Hash, Eq, PartialEq, Clone)]
138pub enum NetlinkFamily {
139    Invalid,
140    Route,
141    Usersock,
142    Firewall,
143    SockDiag,
144    Nflog,
145    Xfrm,
146    Selinux,
147    Iscsi,
148    Audit,
149    FibLookup,
150    Connector,
151    Netfilter,
152    Ip6Fw,
153    Dnrtmsg,
154    KobjectUevent,
155    Generic,
156    Scsitransport,
157    Ecryptfs,
158    Rdma,
159    Crypto,
160    Smc,
161}
162
163impl NetlinkFamily {
164    pub fn from_raw(family: u32) -> Self {
165        match family {
166            NETLINK_ROUTE => NetlinkFamily::Route,
167            NETLINK_USERSOCK => NetlinkFamily::Usersock,
168            NETLINK_FIREWALL => NetlinkFamily::Firewall,
169            NETLINK_SOCK_DIAG => NetlinkFamily::SockDiag,
170            NETLINK_NFLOG => NetlinkFamily::Nflog,
171            NETLINK_XFRM => NetlinkFamily::Xfrm,
172            NETLINK_SELINUX => NetlinkFamily::Selinux,
173            NETLINK_ISCSI => NetlinkFamily::Iscsi,
174            NETLINK_AUDIT => NetlinkFamily::Audit,
175            NETLINK_FIB_LOOKUP => NetlinkFamily::FibLookup,
176            NETLINK_CONNECTOR => NetlinkFamily::Connector,
177            NETLINK_NETFILTER => NetlinkFamily::Netfilter,
178            NETLINK_IP6_FW => NetlinkFamily::Ip6Fw,
179            NETLINK_DNRTMSG => NetlinkFamily::Dnrtmsg,
180            NETLINK_KOBJECT_UEVENT => NetlinkFamily::KobjectUevent,
181            NETLINK_GENERIC => NetlinkFamily::Generic,
182            NETLINK_SCSITRANSPORT => NetlinkFamily::Scsitransport,
183            NETLINK_ECRYPTFS => NetlinkFamily::Ecryptfs,
184            NETLINK_RDMA => NetlinkFamily::Rdma,
185            NETLINK_CRYPTO => NetlinkFamily::Crypto,
186            NETLINK_SMC => NetlinkFamily::Smc,
187            _ => NetlinkFamily::Invalid,
188        }
189    }
190
191    pub fn as_raw(&self) -> u32 {
192        match self {
193            NetlinkFamily::Route => NETLINK_ROUTE,
194            NetlinkFamily::KobjectUevent => NETLINK_KOBJECT_UEVENT,
195            NetlinkFamily::Audit => NETLINK_AUDIT,
196            _ => 0,
197        }
198    }
199}
200
201struct NetlinkSocketInner {
202    /// The specific type of netlink socket.
203    family: NetlinkFamily,
204
205    /// The [`MessageQueue`] that contains messages from netlink to the client.
206    receive_buffer: MessageQueue,
207
208    /// The socket's send buffer size. Note, This value is only used
209    /// to serve getsockopt calls for `SO_SNDBUF`. It does not yet enforce a
210    /// limit on the number of messages netlink will buffer from the client.
211    /// TODO(https://fxbug.dev/285880057): Limit the size of the send buffer.
212    send_buf_size: usize,
213
214    /// This queue will be notified on reads, writes, disconnects etc.
215    waiters: WaitQueue,
216
217    /// The address of this socket.
218    address: Option<NetlinkAddress>,
219
220    /// See SO_PASSCRED.
221    pub passcred: bool,
222
223    /// See SO_TIMESTAMP.
224    pub timestamp: bool,
225
226    /// See NETLINK_GET_STRICT_CHK.
227    pub strict_chk: bool,
228}
229
230impl NetlinkSocketInner {
231    fn new(family: NetlinkFamily) -> Self {
232        Self {
233            family,
234            receive_buffer: MessageQueue::new(SOCKET_DEFAULT_SIZE),
235            send_buf_size: SOCKET_DEFAULT_SIZE,
236            waiters: WaitQueue::default(),
237            address: None,
238            passcred: false,
239            timestamp: false,
240            strict_chk: false,
241        }
242    }
243
244    fn bind(
245        &mut self,
246        current_task: &CurrentTask,
247        socket_address: SocketAddress,
248    ) -> Result<(), Errno> {
249        if self.address.is_some() {
250            return error!(EINVAL);
251        }
252
253        let netlink_address = match socket_address {
254            SocketAddress::Netlink(mut netlink_address) => {
255                // TODO: Support distinct IDs for processes with multiple netlink sockets.
256                netlink_address.set_pid_if_zero(current_task.get_pid());
257                netlink_address
258            }
259            _ => return error!(EINVAL),
260        };
261
262        self.address = Some(netlink_address);
263        Ok(())
264    }
265
266    fn connect(&mut self, current_task: &CurrentTask, peer: SocketPeer) -> Result<(), Errno> {
267        let address = match peer {
268            SocketPeer::Address(address) => address,
269            _ => return error!(EINVAL),
270        };
271        // Connect is equivalent to bind, but error are ignored.
272        let _ = self.bind(current_task, address);
273        Ok(())
274    }
275
276    fn read_message(&mut self) -> Option<Message> {
277        let message = self.receive_buffer.read_message();
278        if message.is_some() {
279            self.waiters.notify_fd_events(FdEvents::POLLOUT);
280        }
281        message
282    }
283
284    fn read_datagram(
285        &mut self,
286        data: &mut dyn OutputBuffer,
287        flags: SocketMessageFlags,
288    ) -> Result<MessageReadInfo, Errno> {
289        let mut info = if flags.contains(SocketMessageFlags::PEEK) {
290            self.receive_buffer.peek_datagram(data)
291        } else {
292            self.receive_buffer.read_datagram(data)
293        }?;
294        if info.message_length == 0 {
295            return error!(EAGAIN);
296        }
297
298        if self.passcred {
299            track_stub!(TODO("https://fxbug.dev/297373991"), "SCM_CREDENTIALS/SO_PASSCRED");
300            info.ancillary_data.push(AncillaryData::Unix(UnixControlData::unknown_creds()));
301        }
302
303        Ok(info)
304    }
305
306    fn write_to_queue(
307        &mut self,
308        data: &mut dyn InputBuffer,
309        address: Option<NetlinkAddress>,
310        ancillary_data: &mut Vec<AncillaryData>,
311    ) -> Result<usize, Errno> {
312        let socket_address = match address {
313            Some(addr) => Some(SocketAddress::Netlink(addr)),
314            None => self.address.as_ref().map(|addr| SocketAddress::Netlink(addr.clone())),
315        };
316        let bytes_written =
317            self.receive_buffer.write_datagram(data, socket_address, ancillary_data)?;
318        if bytes_written > 0 {
319            self.waiters.notify_fd_events(FdEvents::POLLIN);
320        }
321        Ok(bytes_written)
322    }
323
324    fn wait_async(
325        &mut self,
326        waiter: &Waiter,
327        events: FdEvents,
328        handler: EventHandler,
329    ) -> WaitCanceler {
330        self.waiters.wait_async_fd_events(waiter, events, handler)
331    }
332
333    fn query_events(&self) -> FdEvents {
334        self.receive_buffer.query_events()
335    }
336
337    fn getsockname(&self) -> Result<SocketAddress, Errno> {
338        match &self.address {
339            Some(addr) => Ok(SocketAddress::Netlink(addr.clone())),
340            _ => Ok(SocketAddress::default_for_domain(SocketDomain::Netlink)),
341        }
342    }
343
344    fn getpeername(&self) -> Result<SocketAddress, Errno> {
345        match &self.address {
346            Some(addr) => Ok(SocketAddress::Netlink(addr.clone())),
347            _ => Ok(SocketAddress::default_for_domain(SocketDomain::Netlink)),
348        }
349    }
350
351    fn getsockopt(&self, level: u32, optname: u32) -> Result<Vec<u8>, Errno> {
352        let opt_value = match level {
353            SOL_SOCKET => match optname {
354                SO_PASSCRED => (self.passcred as u32).as_bytes().to_vec(),
355                SO_TIMESTAMP => (self.timestamp as u32).as_bytes().to_vec(),
356                SO_SNDBUF => (self.send_buf_size as socklen_t).to_ne_bytes().to_vec(),
357                SO_RCVBUF => (self.receive_buffer.capacity() as socklen_t).to_ne_bytes().to_vec(),
358                SO_SNDBUFFORCE => (self.send_buf_size as socklen_t).to_ne_bytes().to_vec(),
359                SO_RCVBUFFORCE => {
360                    (self.receive_buffer.capacity() as socklen_t).to_ne_bytes().to_vec()
361                }
362                SO_PROTOCOL => self.family.as_raw().as_bytes().to_vec(),
363                _ => return error!(ENOSYS),
364            },
365            SOL_NETLINK => match optname {
366                NETLINK_GET_STRICT_CHK => (self.strict_chk as u32).as_bytes().to_vec(),
367                _ => return error!(ENOSYS),
368            },
369            _ => vec![],
370        };
371
372        Ok(opt_value)
373    }
374
375    fn setsockopt(
376        &mut self,
377        current_task: &CurrentTask,
378        level: u32,
379        optname: u32,
380        optval: SockOptValue,
381    ) -> Result<(), Errno> {
382        match level {
383            SOL_SOCKET => match optname {
384                SO_SNDBUF => {
385                    let requested_capacity: socklen_t = optval.read(current_task)?;
386                    // SO_SNDBUF doubles the requested capacity to leave space for bookkeeping.
387                    // See https://man7.org/linux/man-pages/man7/socket.7.html
388                    let capacity = usize::try_from(requested_capacity * 2).unwrap_or(usize::MAX);
389                    // TODO(https://fxbug.dev/322907334): Clamp to `wmem_max`.
390                    let capacity = capacity.clamp(SOCKET_MIN_SIZE, SOCKET_MAX_SIZE);
391                    self.send_buf_size = capacity;
392                }
393                SO_SNDBUFFORCE => {
394                    security::check_task_capable(current_task, CAP_NET_ADMIN)?;
395                    let requested_capacity: socklen_t = optval.read(current_task)?;
396                    // SO_SNDBUFFORE doubles the requested capacity to leave space for bookkeeping.
397                    // See https://man7.org/linux/man-pages/man7/socket.7.html
398                    let capacity = usize::try_from(requested_capacity * 2).unwrap_or(usize::MAX);
399                    self.send_buf_size = capacity;
400                }
401                SO_RCVBUF => {
402                    let requested_capacity: socklen_t = optval.read(current_task)?;
403                    // SO_RCVBUF doubles the requested capacity to leave space for bookkeeping.
404                    // See https://man7.org/linux/man-pages/man7/socket.7.html
405                    let capacity = usize::try_from(requested_capacity * 2).unwrap_or(usize::MAX);
406                    // TODO(https://fxbug.dev/322906968): Clamp to `rmem_max`.
407                    let capacity = capacity.clamp(SOCKET_MIN_SIZE, SOCKET_MAX_SIZE);
408                    self.receive_buffer.set_capacity(capacity)?;
409                }
410                SO_RCVBUFFORCE => {
411                    security::check_task_capable(current_task, CAP_NET_ADMIN)?;
412                    let requested_capacity: socklen_t = optval.read(current_task)?;
413                    // SO_RCVBUFFORE doubles the requested capacity to leave space for bookkeeping.
414                    // See https://man7.org/linux/man-pages/man7/socket.7.html
415                    let capacity = usize::try_from(requested_capacity * 2).unwrap_or(usize::MAX);
416                    self.receive_buffer.set_capacity(capacity)?;
417                }
418                SO_PASSCRED => {
419                    let passcred: u32 = optval.read(current_task)?;
420                    self.passcred = passcred != 0;
421                }
422                SO_TIMESTAMP => {
423                    let timestamp: u32 = optval.read(current_task)?;
424                    self.timestamp = timestamp != 0;
425                }
426                _ => return error!(ENOSYS),
427            },
428            SOL_NETLINK => match optname {
429                NETLINK_GET_STRICT_CHK => {
430                    let strict_chk: u32 = optval.read(current_task)?;
431                    self.strict_chk = strict_chk != 0;
432                }
433                _ => return error!(ENOSYS),
434            },
435            _ => return error!(ENOSYS),
436        }
437
438        Ok(())
439    }
440}
441
442/// A fake Netlink socket that loops messages back to the client.
443///
444/// Used as a placeholder implementation for protocol families that lack a real
445/// implementation.
446struct StubbedNetlinkSocket {
447    inner: LockDepMutex<NetlinkSocketInner, NetlinkSocketInnerLock>,
448}
449
450impl StubbedNetlinkSocket {
451    pub fn new(family: NetlinkFamily) -> Self {
452        track_stub!(
453            TODO("https://fxbug.dev/278565021"),
454            format!("Creating StubbedNetlinkSocket: {:?}", family).as_str()
455        );
456        StubbedNetlinkSocket { inner: NetlinkSocketInner::new(family).into() }
457    }
458
459    /// Locks and returns the inner state of the Socket.
460    fn lock(&self) -> LockDepGuard<'_, NetlinkSocketInner> {
461        self.inner.lock()
462    }
463}
464
465impl SocketOps for StubbedNetlinkSocket {
466    fn connect(
467        &self,
468        _socket: &SocketHandle,
469        current_task: &CurrentTask,
470        peer: SocketPeer,
471    ) -> Result<(), Errno> {
472        self.lock().connect(current_task, peer)
473    }
474
475    fn listen(&self, _socket: &Socket, _backlog: i32, _credentials: ucred) -> Result<(), Errno> {
476        error!(EOPNOTSUPP)
477    }
478
479    fn accept(&self, _socket: &Socket, _current_task: &CurrentTask) -> Result<SocketHandle, Errno> {
480        error!(EOPNOTSUPP)
481    }
482
483    fn bind(
484        &self,
485        _socket: &Socket,
486        current_task: &CurrentTask,
487        socket_address: SocketAddress,
488    ) -> Result<(), Errno> {
489        self.lock().bind(current_task, socket_address)
490    }
491
492    fn read(
493        &self,
494        _socket: &Socket,
495        _current_task: &CurrentTask,
496        data: &mut dyn OutputBuffer,
497        _flags: SocketMessageFlags,
498    ) -> Result<MessageReadInfo, Errno> {
499        let msg = self.lock().read_message();
500        match msg {
501            Some(message) => {
502                // Mark the message as complete and return it.
503                let (mut nl_msg, _) =
504                    nlmsghdr::read_from_prefix(&message.data).map_err(|_| errno!(EINVAL))?;
505                nl_msg.nlmsg_type = NLMSG_DONE as u16;
506                nl_msg.nlmsg_flags &= NLM_F_MULTI as u16;
507                let msg_bytes = nl_msg.as_bytes();
508                let bytes_read = data.write(msg_bytes)?;
509
510                let info = MessageReadInfo {
511                    bytes_read,
512                    message_length: msg_bytes.len(),
513                    address: Some(SocketAddress::Netlink(NetlinkAddress::default())),
514                    ancillary_data: vec![],
515                };
516                Ok(info)
517            }
518            None => Ok(MessageReadInfo::default()),
519        }
520    }
521
522    fn write(
523        &self,
524        _socket: &Socket,
525        _current_task: &CurrentTask,
526        data: &mut dyn InputBuffer,
527        dest_address: &mut Option<SocketAddress>,
528        ancillary_data: &mut Vec<AncillaryData>,
529    ) -> Result<usize, Errno> {
530        let mut local_address = self.lock().address.clone();
531
532        let destination = match dest_address {
533            Some(SocketAddress::Netlink(addr)) => addr,
534            _ => match &mut local_address {
535                Some(addr) => addr,
536                _ => return Ok(data.drain()),
537            },
538        };
539
540        if destination.groups != 0 {
541            track_stub!(TODO("https://fxbug.dev/322874956"), "StubbedNetlinkSockets multicasting");
542            return Ok(data.drain());
543        }
544
545        self.lock().write_to_queue(data, Some(NetlinkAddress::default()), ancillary_data)
546    }
547
548    fn wait_async(
549        &self,
550        _socket: &Socket,
551        _current_task: &CurrentTask,
552        waiter: &Waiter,
553        events: FdEvents,
554        handler: EventHandler,
555    ) -> WaitCanceler {
556        self.lock().wait_async(waiter, events, handler)
557    }
558
559    fn query_events(
560        &self,
561        _socket: &Socket,
562        _current_task: &CurrentTask,
563    ) -> Result<FdEvents, Errno> {
564        Ok(self.lock().query_events() & FdEvents::POLLIN)
565    }
566
567    fn shutdown(&self, _socket: &Socket, _how: SocketShutdownFlags) -> Result<(), Errno> {
568        track_stub!(TODO("https://fxbug.dev/322875507"), "StubbedNetlinkSocket::shutdown");
569        Ok(())
570    }
571
572    fn close(&self, _current_task: &CurrentTask, _socket: &Socket) {}
573
574    fn getsockname(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
575        self.lock().getsockname()
576    }
577
578    fn getpeername(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
579        self.lock().getpeername()
580    }
581
582    fn getsockopt(
583        &self,
584        _socket: &Socket,
585        _current_task: &CurrentTask,
586        level: u32,
587        optname: u32,
588        _optlen: u32,
589    ) -> Result<Vec<u8>, Errno> {
590        self.lock().getsockopt(level, optname)
591    }
592
593    fn setsockopt(
594        &self,
595        _socket: &Socket,
596        current_task: &CurrentTask,
597        level: u32,
598        optname: u32,
599        optval: SockOptValue,
600    ) -> Result<(), Errno> {
601        self.lock().setsockopt(current_task, level, optname, optval)
602    }
603}
604
605/// Socket implementation for the NETLINK_KOBJECT_UEVENT family of netlink sockets.
606struct UEventNetlinkSocket {
607    inner: Arc<LockDepMutex<NetlinkSocketInner, NetlinkSocketInnerLock>>,
608    device_listener_key:
609        LockDepMutex<Option<DeviceListenerKey>, UEventNetlinkSocketDeviceListenerKeyLock>,
610}
611
612impl Default for UEventNetlinkSocket {
613    #[allow(clippy::let_and_return)]
614    fn default() -> Self {
615        let result = Self {
616            inner: Arc::new(NetlinkSocketInner::new(NetlinkFamily::KobjectUevent).into()),
617            device_listener_key: Default::default(),
618        };
619        #[cfg(any(test, debug_assertions))]
620        {
621            let _l1 = result.device_listener_key.lock();
622            let _l2 = result.lock();
623        }
624        result
625    }
626}
627
628impl UEventNetlinkSocket {
629    /// Locks and returns the inner state of the Socket.
630    fn lock(&self) -> LockDepGuard<'_, NetlinkSocketInner> {
631        self.inner.lock()
632    }
633
634    fn register_listener(
635        &self,
636        current_task: &CurrentTask,
637        state: LockDepGuard<'_, NetlinkSocketInner>,
638    ) {
639        if state.address.is_none() {
640            return;
641        }
642        std::mem::drop(state);
643        let mut key_state = self.device_listener_key.lock();
644        if key_state.is_none() {
645            *key_state =
646                Some(current_task.kernel().device_registry.register_listener(self.inner.clone()));
647        }
648    }
649}
650
651impl SocketOps for UEventNetlinkSocket {
652    fn connect(
653        &self,
654        _socket: &SocketHandle,
655        current_task: &CurrentTask,
656        peer: SocketPeer,
657    ) -> Result<(), Errno> {
658        let mut state = self.lock();
659        state.connect(current_task, peer)?;
660        self.register_listener(current_task, state);
661        Ok(())
662    }
663
664    fn listen(&self, _socket: &Socket, _backlog: i32, _credentials: ucred) -> Result<(), Errno> {
665        error!(EOPNOTSUPP)
666    }
667
668    fn accept(&self, _socket: &Socket, _current_task: &CurrentTask) -> Result<SocketHandle, Errno> {
669        error!(EOPNOTSUPP)
670    }
671
672    fn bind(
673        &self,
674        _socket: &Socket,
675        current_task: &CurrentTask,
676        socket_address: SocketAddress,
677    ) -> Result<(), Errno> {
678        let mut state = self.lock();
679        state.bind(current_task, socket_address)?;
680        self.register_listener(current_task, state);
681        Ok(())
682    }
683
684    fn read(
685        &self,
686        _socket: &Socket,
687        _current_task: &CurrentTask,
688        data: &mut dyn OutputBuffer,
689        flags: SocketMessageFlags,
690    ) -> Result<MessageReadInfo, Errno> {
691        self.lock().read_datagram(data, flags)
692    }
693
694    fn write(
695        &self,
696        _socket: &Socket,
697        _current_task: &CurrentTask,
698        _data: &mut dyn InputBuffer,
699        _dest_address: &mut Option<SocketAddress>,
700        _ancillary_data: &mut Vec<AncillaryData>,
701    ) -> Result<usize, Errno> {
702        error!(EOPNOTSUPP)
703    }
704
705    fn wait_async(
706        &self,
707        _socket: &Socket,
708        _current_task: &CurrentTask,
709        waiter: &Waiter,
710        events: FdEvents,
711        handler: EventHandler,
712    ) -> WaitCanceler {
713        self.lock().wait_async(waiter, events, handler)
714    }
715
716    fn query_events(
717        &self,
718        _socket: &Socket,
719        _current_task: &CurrentTask,
720    ) -> Result<FdEvents, Errno> {
721        Ok(self.lock().query_events() & FdEvents::POLLIN)
722    }
723
724    fn shutdown(&self, _socket: &Socket, _how: SocketShutdownFlags) -> Result<(), Errno> {
725        track_stub!(TODO("https://fxbug.dev/322875507"), "UEventNetlinkSocket::shutdown");
726        Ok(())
727    }
728
729    fn close(&self, current_task: &CurrentTask, _socket: &Socket) {
730        let id = self.device_listener_key.lock().take();
731        if let Some(id) = id {
732            current_task.kernel().device_registry.unregister_listener(&id);
733        }
734    }
735
736    fn getsockname(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
737        self.lock().getsockname()
738    }
739
740    fn getpeername(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
741        self.lock().getpeername()
742    }
743
744    fn getsockopt(
745        &self,
746        _socket: &Socket,
747        _current_task: &CurrentTask,
748        level: u32,
749        optname: u32,
750        _optlen: u32,
751    ) -> Result<Vec<u8>, Errno> {
752        self.lock().getsockopt(level, optname)
753    }
754
755    fn setsockopt(
756        &self,
757        _socket: &Socket,
758        current_task: &CurrentTask,
759        level: u32,
760        optname: u32,
761        optval: SockOptValue,
762    ) -> Result<(), Errno> {
763        self.lock().setsockopt(current_task, level, optname, optval)
764    }
765}
766
767impl DeviceListener for Arc<LockDepMutex<NetlinkSocketInner, NetlinkSocketInnerLock>> {
768    fn on_device_event(&self, action: UEventAction, device: Device, context: UEventContext) {
769        let path = device.path_from_depth(0);
770
771        let mut props = device.get_uevent_properties_list();
772
773        // Prepend ACTION and SEQNUM to maintain existing order
774        props.insert(0, (b"ACTION".into(), action.to_string().into()));
775        props.insert(1, (b"SEQNUM".into(), context.seqnum.to_string().into()));
776
777        let flattened = flatten_uevent_properties(props, '\0');
778
779        let mut message = vec![];
780        write!(&mut message, "{action}@/{path}\0", action = action, path = path).unwrap();
781        message.extend_from_slice(flattened.as_ref());
782
783        let ancillary_data = AncillaryData::Unix(UnixControlData::Credentials(Default::default()));
784        let mut ancillary_data = vec![ancillary_data];
785        // Ignore write errors
786        let _ = self.lock().write_to_queue(
787            &mut VecInputBuffer::new(&message),
788            Some(NetlinkAddress { pid: 0, groups: 1 }),
789            &mut ancillary_data,
790        );
791    }
792}
793
794/// Type for sending messages from [`netlink::Netlink`] to an individual socket.
795#[derive(Clone)]
796pub struct NetlinkToClientSender<M> {
797    /// The inner socket implementation, which holds a message queue.
798    inner: Arc<LockDepMutex<NetlinkSocketInner, NetlinkSocketInnerLock>>,
799
800    /// `PhantomData<fn(M) -> M>` is used instead of `PhantomData<M>` in order
801    /// to ensure that the type is invariant over `M` and that it implements
802    /// `Sync` even if `M` is not `Sync`.
803    _message_type: PhantomData<fn(M) -> M>,
804}
805
806impl<M> NetlinkToClientSender<M> {
807    fn new(inner: Arc<LockDepMutex<NetlinkSocketInner, NetlinkSocketInnerLock>>) -> Self {
808        NetlinkToClientSender { _message_type: Default::default(), inner }
809    }
810}
811
812impl<M: Clone + NetlinkSerializable + Send> Sender<M> for NetlinkToClientSender<M> {
813    fn send(&mut self, message: NetlinkMessage<M>, group: Option<ModernGroup>) {
814        // Serialize the message
815        let mut buf = vec![0; message.buffer_len()];
816        message.emit(&mut buf);
817        let mut buf: VecInputBuffer = buf.into();
818        // Write the message into the inner socket buffer.
819        let NetlinkToClientSender { _message_type: _, inner } = self;
820        let mut guard = inner.lock();
821
822        // To avoid dropping messages when the receive buffer is
823        // full, grow the buffer on behalf of the client.
824        // This is a stop gap measure to avoid dropping messages
825        // when netlink produces a large response to a
826        // NLM_F_DUMP request.
827        //
828        // TODO(https://fxbug.dev/459883760): The memory
829        // implications of this may be problematic. It should be
830        // replaced with a proper mechanism to handle a backlog
831        // of NLM_F_DUMP responses.
832        let available = guard.receive_buffer.available_capacity();
833        let required = buf.available();
834        if available < required {
835            let delta = required - available;
836            let current_capacity = guard.receive_buffer.capacity();
837            let new_capacity = (current_capacity + delta).min(SOCKET_MAX_SIZE);
838            match guard.receive_buffer.set_capacity(new_capacity) {
839                Ok(()) => {}
840                Err(e) => {
841                    log_error!(
842                        tag = NETLINK_LOG_TAG;
843                        "Failed to increase receive buffer size: {:?}",
844                        e
845                    );
846                }
847            }
848        }
849
850        let _bytes_written: usize = guard
851            .write_to_queue(
852                &mut buf,
853                Some(NetlinkAddress {
854                    // All messages come from the "kernel" which has PID of 0.
855                    pid: 0,
856                    // If this is a multicast message, set the group the multicast
857                    // message is from.
858                    groups: group
859                        .map(SingleLegacyGroup::try_from)
860                        .and_then(Result::<_, NoMappingFromModernToLegacyGroupError>::ok)
861                        .map_or(0, |g| g.inner()),
862                }),
863                &mut Vec::new(),
864            )
865            .unwrap_or_else(|e| {
866                log_error!(
867                    tag = NETLINK_LOG_TAG;
868                    "Failed to write message into buffer for socket. Errno: {:?}",
869                    e
870                );
871                0
872            });
873    }
874}
875
876#[derive(Clone)]
877pub struct NetlinkAccessControl<'a> {
878    current_task: &'a CurrentTask,
879}
880
881impl<'a> NetlinkAccessControl<'a> {
882    pub fn new(current_task: &'a CurrentTask) -> Self {
883        Self { current_task }
884    }
885}
886
887impl<'a> AccessControl<Arc<Credentials>> for NetlinkAccessControl<'a> {
888    fn grant_assess(
889        &self,
890        creds: &Arc<Credentials>,
891        permission: Permission,
892    ) -> Result<(), netlink::Errno> {
893        let need_cap_net_admin = match permission {
894            Permission::NetlinkRouteRead => false,
895            Permission::NetlinkRouteWrite => true,
896            Permission::NetlinkSockDiagRead => false,
897            Permission::NetlinkSockDiagDestroy => true,
898        };
899        if !need_cap_net_admin {
900            return Ok(());
901        }
902
903        self.current_task.override_creds(creds.clone(), || {
904            security::check_task_capable(self.current_task, CAP_NET_ADMIN).map_err(|error| {
905                netlink::Errno::new(error.code.error_code() as i32)
906                    .expect("Errno::error_code() is expected to be in range [1..max_i32]")
907            })
908        })
909    }
910}
911pub struct NetlinkContextImpl;
912
913impl NetlinkContext for NetlinkContextImpl {
914    type Creds = Arc<Credentials>;
915    type Sender<M: Clone + NetlinkSerializable + Send> = NetlinkToClientSender<M>;
916    type Receiver<
917        M: Send + MessageWithPermission + NetlinkDeserializable<Error: Into<DecodeError>>,
918    > = UnboundedReceiver<NetlinkMessageWithCreds<UnparsedNetlinkMessage<Vec<u8>, M>, Self::Creds>>;
919    type AccessControl<'a> = NetlinkAccessControl<'a>;
920}
921
922fn new_route_socket(kernel: &Arc<Kernel>) -> Result<NetlinkSocket<NetlinkRouteClient>, Errno> {
923    let inner = Arc::new(LockDepMutex::new(NetlinkSocketInner::new(NetlinkFamily::Route)));
924    let (message_sender, message_receiver) = mpsc::unbounded();
925    let client = match kernel
926        .network_netlink()
927        .new_route_client(NetlinkToClientSender::new(inner.clone()), message_receiver)
928    {
929        Ok(client) => client,
930        Err(NewClientError::Disconnected) => {
931            log_error!(
932                tag = NETLINK_LOG_TAG;
933                "Netlink async worker is unexpectedly disconnected"
934            );
935            return error!(EPIPE);
936        }
937    };
938    Ok(NetlinkSocket { inner, client, message_sender })
939}
940
941fn new_sock_diag_socket(
942    kernel: &Arc<Kernel>,
943) -> Result<NetlinkSocket<NetlinkSockDiagClient>, Errno> {
944    let inner = Arc::new(LockDepMutex::new(NetlinkSocketInner::new(NetlinkFamily::SockDiag)));
945    let (message_sender, message_receiver) = mpsc::unbounded();
946    let client = match kernel
947        .network_netlink()
948        .new_sock_diag_client(NetlinkToClientSender::new(inner.clone()), message_receiver)
949    {
950        Ok(client) => client,
951        Err(NewClientError::Disconnected) => {
952            log_error!(
953                tag = NETLINK_LOG_TAG;
954                "Netlink async worker is unexpectedly disconnected"
955            );
956            return error!(EPIPE);
957        }
958    };
959    Ok(NetlinkSocket { inner, client, message_sender })
960}
961
962/// An abstraction over common networking-specific netlink sockets.
963struct NetlinkSocket<C: NetlinkClient> {
964    /// The inner Netlink socket implementation
965    inner: Arc<LockDepMutex<NetlinkSocketInner, NetlinkSocketInnerLock>>,
966    /// The implementation of a client (socket connection) to a netlink protocol
967    /// family.
968    client: C,
969    /// The sender of messages from this socket to Netlink.
970    // TODO(https://issuetracker.google.com/285880057): Bound the capacity of
971    // the "send buffer".
972    message_sender: UnboundedSender<
973        NetlinkMessageWithCreds<UnparsedNetlinkMessage<Vec<u8>, C::Request>, Arc<Credentials>>,
974    >,
975}
976
977/// A type that provides Netlink message deserialization options.
978trait DeserializeOptionsProvider {
979    /// The type of the message to deserialize.
980    type Message: NetlinkDeserializable;
981    /// The options to use when deserializing a `Message`.
982    fn options(&self) -> <Self::Message as NetlinkDeserializable>::DeserializeOptions;
983}
984
985impl DeserializeOptionsProvider for NetlinkSocket<NetlinkRouteClient> {
986    type Message = RouteNetlinkMessage;
987    fn options(&self) -> RouteNetlinkMessageParseMode {
988        let strict = self.inner.lock().strict_chk;
989        if strict {
990            RouteNetlinkMessageParseMode::Strict
991        } else {
992            RouteNetlinkMessageParseMode::Relaxed
993        }
994    }
995}
996
997impl DeserializeOptionsProvider for NetlinkSocket<NetlinkSockDiagClient> {
998    type Message = SockDiagRequest;
999    fn options(&self) -> EmptyDeserializeSockDiagOptions {
1000        EmptyDeserializeSockDiagOptions
1001    }
1002}
1003
1004impl<C: NetlinkClient + 'static> SocketOps for NetlinkSocket<C>
1005where
1006    Self: DeserializeOptionsProvider<Message = C::Request>,
1007{
1008    fn connect(
1009        &self,
1010        _socket: &SocketHandle,
1011        current_task: &CurrentTask,
1012        peer: SocketPeer,
1013    ) -> Result<(), Errno> {
1014        let NetlinkSocket { inner, client: _, message_sender: _ } = self;
1015        inner.lock().connect(current_task, peer)
1016    }
1017
1018    fn listen(&self, _socket: &Socket, _backlog: i32, _credentials: ucred) -> Result<(), Errno> {
1019        error!(EOPNOTSUPP)
1020    }
1021
1022    fn accept(&self, _socket: &Socket, _current_task: &CurrentTask) -> Result<SocketHandle, Errno> {
1023        error!(EOPNOTSUPP)
1024    }
1025
1026    fn bind(
1027        &self,
1028        _socket: &Socket,
1029        current_task: &CurrentTask,
1030        socket_address: SocketAddress,
1031    ) -> Result<(), Errno> {
1032        let NetlinkSocket { inner, client, message_sender: _ } = self;
1033
1034        let multicast_groups = match &socket_address {
1035            SocketAddress::Netlink(NetlinkAddress { pid: _, groups }) => *groups,
1036            _ => return error!(EINVAL),
1037        };
1038        let pid = {
1039            let mut inner = inner.lock();
1040            inner.bind(current_task, socket_address)?;
1041            inner
1042                .address
1043                .as_ref()
1044                .and_then(|NetlinkAddress { pid, groups: _ }| NonZeroU32::new(*pid))
1045        };
1046        if let Some(pid) = pid {
1047            client.set_pid(pid);
1048        }
1049        // This "blocks" in order to synchronize with the internal
1050        // state of the netlink worker, but we're not blocking on
1051        // the completion of any i/o or any expensive computation,
1052        // so there's no need to support interrupts here.
1053        client
1054            .set_legacy_memberships(LegacyGroups(multicast_groups))
1055            .map_err(|InvalidLegacyGroupsError {}| errno!(EPERM))?
1056            .wait_until_complete();
1057        Ok(())
1058    }
1059
1060    fn read(
1061        &self,
1062        _socket: &Socket,
1063        _current_task: &CurrentTask,
1064        data: &mut dyn OutputBuffer,
1065        flags: SocketMessageFlags,
1066    ) -> Result<MessageReadInfo, Errno> {
1067        let NetlinkSocket { inner, client: _, message_sender: _ } = self;
1068        inner.lock().read_datagram(data, flags)
1069    }
1070
1071    fn write(
1072        &self,
1073        socket: &Socket,
1074        current_task: &CurrentTask,
1075        data: &mut dyn InputBuffer,
1076        _dest_address: &mut Option<SocketAddress>,
1077        _ancillary_data: &mut Vec<AncillaryData>,
1078    ) -> Result<usize, Errno> {
1079        let NetlinkSocket { inner: _, client: _, message_sender } = self;
1080
1081        let bytes = data.peek_all()?;
1082        let bytes_len = bytes.len();
1083
1084        // Parse only the netlink header to send it through security check.
1085        match NetlinkBuffer::new(&bytes) {
1086            Ok(buffer) => {
1087                security::check_netlink_send_access(current_task, socket, buffer.message_type())?;
1088            }
1089            Err(e) => {
1090                // If we can't even decode the header of the netlink message,
1091                // then return early here as a stronger statement that we're not
1092                // going to accidentally operate on it and violate the security
1093                // check. The netlink crate would end up dropping this with no
1094                // response as well.
1095                log_warn!(tag = NETLINK_LOG_TAG;
1096                    "Failed to parse netlink header {e:?}"
1097                );
1098                data.drain();
1099                return Ok(bytes_len);
1100            }
1101        }
1102
1103        let msg = NetlinkMessageWithCreds::new(
1104            UnparsedNetlinkMessage::new(bytes, self.options()),
1105            current_task.current_creds().clone(),
1106        );
1107        message_sender.unbounded_send(msg).map_err(|e| {
1108            log_warn!(
1109                tag = NETLINK_LOG_TAG;
1110                "Netlink receiver unexpectedly disconnected for socket: {:?}",
1111                e
1112            );
1113            errno!(EPIPE)
1114        })?;
1115        data.drain();
1116        Ok(bytes_len)
1117    }
1118
1119    fn wait_async(
1120        &self,
1121        _socket: &Socket,
1122        _current_task: &CurrentTask,
1123        waiter: &Waiter,
1124        events: FdEvents,
1125        handler: EventHandler,
1126    ) -> WaitCanceler {
1127        let NetlinkSocket { inner, client: _, message_sender: _ } = self;
1128        inner.lock().wait_async(waiter, events, handler)
1129    }
1130
1131    fn query_events(
1132        &self,
1133        _socket: &Socket,
1134        _current_task: &CurrentTask,
1135    ) -> Result<FdEvents, Errno> {
1136        let NetlinkSocket { inner, client: _, message_sender: _ } = self;
1137        Ok(inner.lock().query_events() & FdEvents::POLLIN)
1138    }
1139
1140    fn shutdown(&self, _socket: &Socket, _how: SocketShutdownFlags) -> Result<(), Errno> {
1141        error!(EOPNOTSUPP)
1142    }
1143
1144    fn close(&self, _current_task: &CurrentTask, _socket: &Socket) {
1145        // Close the underlying channel to the Netlink worker.
1146        self.message_sender.close_channel();
1147    }
1148
1149    fn getsockname(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
1150        let NetlinkSocket { inner, client: _, message_sender: _ } = self;
1151        inner.lock().getsockname()
1152    }
1153
1154    fn getpeername(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
1155        self.inner.lock().getpeername()
1156    }
1157
1158    fn getsockopt(
1159        &self,
1160        _socket: &Socket,
1161        _current_task: &CurrentTask,
1162        level: u32,
1163        optname: u32,
1164        _optlen: u32,
1165    ) -> Result<Vec<u8>, Errno> {
1166        self.inner.lock().getsockopt(level, optname)
1167    }
1168
1169    fn setsockopt(
1170        &self,
1171        _socket: &Socket,
1172        current_task: &CurrentTask,
1173        level: u32,
1174        optname: u32,
1175        optval: SockOptValue,
1176    ) -> Result<(), Errno> {
1177        match (level, optname) {
1178            (SOL_NETLINK, NETLINK_ADD_MEMBERSHIP) => {
1179                let NetlinkSocket { inner: _, client, message_sender: _ } = self;
1180                let group: u32 = optval.read(current_task)?;
1181                let async_work = client
1182                    .add_membership(ModernGroup(group))
1183                    .map_err(|InvalidModernGroupError| errno!(EINVAL))?;
1184                // This "blocks" in order to synchronize with the internal
1185                // state of the rtnetlink worker, but we're not blocking on
1186                // the completion of any i/o or any expensive computation,
1187                // so there's no need to support interrupts here.
1188                async_work.wait_until_complete();
1189                Ok(())
1190            }
1191            (SOL_NETLINK, NETLINK_DROP_MEMBERSHIP) => {
1192                let NetlinkSocket { inner: _, client, message_sender: _ } = self;
1193                let group: u32 = optval.read(current_task)?;
1194                client
1195                    .del_membership(ModernGroup(group))
1196                    .map_err(|InvalidModernGroupError| errno!(EINVAL))?;
1197                Ok(())
1198            }
1199            _ => self.inner.lock().setsockopt(current_task, level, optname, optval),
1200        }
1201    }
1202}
1203
1204/// Socket implementation for the NETLINK_GENERIC family of netlink sockets.
1205struct GenericNetlinkSocket {
1206    inner: Arc<LockDepMutex<NetlinkSocketInner, NetlinkSocketInnerLock>>,
1207    client: GenericNetlinkClientHandle<NetlinkToClientSender<GenericMessage>>,
1208    message_sender: mpsc::UnboundedSender<NetlinkMessage<GenericMessage>>,
1209}
1210
1211impl GenericNetlinkSocket {
1212    pub fn new(kernel: &Kernel) -> Result<Self, Errno> {
1213        let inner = Arc::new(LockDepMutex::new(NetlinkSocketInner::new(NetlinkFamily::Generic)));
1214        let (message_sender, message_receiver) = mpsc::unbounded();
1215        match kernel
1216            .generic_netlink()
1217            .new_generic_client(NetlinkToClientSender::new(inner.clone()), message_receiver)
1218        {
1219            Ok(client) => Ok(Self { inner, client, message_sender }),
1220            Err(e) => {
1221                log_warn!(
1222                    tag = NETLINK_LOG_TAG;
1223                    "Failed to connect to generic netlink server. Errno: {:?}",
1224                    e
1225                );
1226                error!(EPIPE)
1227            }
1228        }
1229    }
1230
1231    /// Locks and returns the inner state of the Socket.
1232    fn lock(&self) -> LockDepGuard<'_, NetlinkSocketInner> {
1233        self.inner.lock()
1234    }
1235}
1236
1237impl SocketOps for GenericNetlinkSocket {
1238    fn connect(
1239        &self,
1240        _socket: &SocketHandle,
1241        current_task: &CurrentTask,
1242        peer: SocketPeer,
1243    ) -> Result<(), Errno> {
1244        let mut state = self.lock();
1245        state.connect(current_task, peer)
1246    }
1247
1248    fn listen(&self, _socket: &Socket, _backlog: i32, _credentials: ucred) -> Result<(), Errno> {
1249        error!(EOPNOTSUPP)
1250    }
1251
1252    fn accept(&self, _socket: &Socket, _current_task: &CurrentTask) -> Result<SocketHandle, Errno> {
1253        error!(EOPNOTSUPP)
1254    }
1255
1256    fn bind(
1257        &self,
1258        _socket: &Socket,
1259        current_task: &CurrentTask,
1260        socket_address: SocketAddress,
1261    ) -> Result<(), Errno> {
1262        let mut state = self.lock();
1263        state.bind(current_task, socket_address)
1264    }
1265
1266    fn read(
1267        &self,
1268        _socket: &Socket,
1269        _current_task: &CurrentTask,
1270        data: &mut dyn OutputBuffer,
1271        flags: SocketMessageFlags,
1272    ) -> Result<MessageReadInfo, Errno> {
1273        self.lock().read_datagram(data, flags)
1274    }
1275
1276    fn write(
1277        &self,
1278        _socket: &Socket,
1279        _current_task: &CurrentTask,
1280        data: &mut dyn InputBuffer,
1281        _dest_address: &mut Option<SocketAddress>,
1282        _ancillary_data: &mut Vec<AncillaryData>,
1283    ) -> Result<usize, Errno> {
1284        let bytes = data.read_all()?;
1285        match NetlinkMessage::<GenericMessage>::deserialize(&bytes, EmptyDeserializeGenlOptions) {
1286            Err(e) => {
1287                log_warn!("Failed to process write; data could not be deserialized: {:?}", e);
1288                error!(EINVAL)
1289            }
1290            Ok(msg) => match self.message_sender.unbounded_send(msg) {
1291                Ok(()) => Ok(bytes.len()),
1292                Err(e) => {
1293                    log_warn!("Netlink receiver unexpectedly disconnected for socket: {:?}", e);
1294                    error!(EPIPE)
1295                }
1296            },
1297        }
1298    }
1299
1300    fn wait_async(
1301        &self,
1302        _socket: &Socket,
1303        _current_task: &CurrentTask,
1304        waiter: &Waiter,
1305        events: FdEvents,
1306        handler: EventHandler,
1307    ) -> WaitCanceler {
1308        self.lock().wait_async(waiter, events, handler)
1309    }
1310
1311    fn query_events(
1312        &self,
1313        _socket: &Socket,
1314        _current_task: &CurrentTask,
1315    ) -> Result<FdEvents, Errno> {
1316        Ok(self.lock().query_events() & FdEvents::POLLIN)
1317    }
1318
1319    fn shutdown(&self, _socket: &Socket, _how: SocketShutdownFlags) -> Result<(), Errno> {
1320        track_stub!(TODO("https://fxbug.dev/322875507"), "GenericNetlinkSocket::shutdown");
1321        Ok(())
1322    }
1323
1324    fn close(&self, _current_task: &CurrentTask, _socket: &Socket) {}
1325
1326    fn getsockname(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
1327        self.lock().getsockname()
1328    }
1329
1330    fn getpeername(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
1331        self.lock().getpeername()
1332    }
1333
1334    fn getsockopt(
1335        &self,
1336        _socket: &Socket,
1337        _current_task: &CurrentTask,
1338        level: u32,
1339        optname: u32,
1340        _optlen: u32,
1341    ) -> Result<Vec<u8>, Errno> {
1342        self.lock().getsockopt(level, optname)
1343    }
1344
1345    fn setsockopt(
1346        &self,
1347        _socket: &Socket,
1348        current_task: &CurrentTask,
1349        level: u32,
1350        optname: u32,
1351        optval: SockOptValue,
1352    ) -> Result<(), Errno> {
1353        match (level, optname) {
1354            (SOL_NETLINK, NETLINK_ADD_MEMBERSHIP) => {
1355                let group_id: u32 = optval.read(current_task)?;
1356                self.client.add_membership(ModernGroup(group_id))
1357            }
1358            _ => self.lock().setsockopt(current_task, level, optname, optval),
1359        }
1360    }
1361}
1362
1363/// Audit client that can be attached to the `AuditLogger`.
1364pub struct AuditNetlinkClient {
1365    /// Reference to the `AuditLogger`.
1366    audit_logger: Arc<AuditLogger>,
1367    /// The waiters queue present in `AuditNetlinkSocket`.
1368    waiters: WaitQueue,
1369    /// Optional response from the `AuditLogger`.
1370    audit_response:
1371        LockDepMutex<Option<NetlinkMessage<GenericMessage>>, AuditNetlinkClientAuditResponseLock>,
1372}
1373
1374impl AuditNetlinkClient {
1375    fn new(audit_logger: Arc<AuditLogger>) -> Self {
1376        Self { audit_logger, waiters: Default::default(), audit_response: Default::default() }
1377    }
1378
1379    pub fn notify(&self) {
1380        self.waiters.notify_fd_events(FdEvents::POLLIN);
1381    }
1382
1383    /// Function to check the capabilities of the current task against CAP_AUDIT_*
1384    fn check_audit_access(
1385        &self,
1386        current_task: &CurrentTask,
1387        request_type: &AuditRequest,
1388    ) -> Result<(), Errno> {
1389        match request_type {
1390            AuditRequest::AuditGet | AuditRequest::AuditSet => {
1391                security::check_task_capable(current_task, CAP_AUDIT_CONTROL)
1392            }
1393            AuditRequest::AuditUser => security::check_task_capable(current_task, CAP_AUDIT_WRITE),
1394        }
1395    }
1396
1397    /// Function to process request coming from userspace, it returns the response after processing
1398    fn process_request(
1399        self: &Arc<Self>,
1400        current_task: &CurrentTask,
1401        nl_message: NetlinkMessage<GenericMessage>,
1402    ) -> Result<NetlinkMessage<GenericMessage>, Errno> {
1403        let (nl_header, nl_payload) = nl_message.into_parts();
1404        let audit_request_type = AuditRequest::try_from(nl_header.message_type as u32)?;
1405        self.check_audit_access(current_task, &audit_request_type)?;
1406
1407        // If there is no GenericMessage, return an ErrorMessage.
1408        let NetlinkPayload::InnerMessage(GenericMessage::Other { payload, .. }) = nl_payload else {
1409            return error!(EINVAL);
1410        };
1411        match audit_request_type {
1412            AuditRequest::AuditGet => self.process_get_status(nl_header.sequence_number),
1413            AuditRequest::AuditSet => self.process_set_status(current_task, nl_header, payload),
1414            AuditRequest::AuditUser => self.process_user_audit(nl_header, payload),
1415        }
1416    }
1417
1418    fn get_nl_response(&self, flags: SocketMessageFlags) -> Option<Vec<u8>> {
1419        if flags.contains(SocketMessageFlags::PEEK) {
1420            if let Some(message) = self.audit_response.lock().as_ref() {
1421                return Some(AuditNetlinkClient::serialize_nlmsg(message.clone()));
1422            }
1423        } else if let Some(message) = self.audit_response.lock().take() {
1424            return Some(AuditNetlinkClient::serialize_nlmsg(message));
1425        }
1426        None
1427    }
1428
1429    /// Function to read an audit message from `AuditLogger`.
1430    fn read_audit_log(self: &Arc<Self>) -> Option<Vec<u8>> {
1431        if let Some(AuditMessage { audit_type, message }) = self.audit_logger.read_audit_log(self) {
1432            return Some(AuditNetlinkClient::serialize_nlmsg(
1433                AuditNetlinkClient::build_audit_nlmsg(0, audit_type, message),
1434            ));
1435        }
1436        None
1437    }
1438
1439    /// Function to read the optional response if present or an audit message.
1440    fn read_nlmsg(self: &Arc<Self>, flags: SocketMessageFlags) -> Result<Vec<u8>, Errno> {
1441        // First check if there is a response and send it if present.
1442        // Send an audit message otherwise or return EAGAIN.
1443        self.get_nl_response(flags).or_else(|| self.read_audit_log()).ok_or_else(|| errno!(EAGAIN))
1444    }
1445
1446    fn process_get_status(
1447        &self,
1448        sequence_number: u32,
1449    ) -> Result<NetlinkMessage<GenericMessage>, Errno> {
1450        Ok(AuditNetlinkClient::build_audit_nlmsg(
1451            sequence_number,
1452            AUDIT_GET as u16,
1453            self.audit_logger.get_status().as_bytes().to_vec(),
1454        ))
1455    }
1456
1457    fn process_set_status(
1458        self: &Arc<Self>,
1459        current_task: &CurrentTask,
1460        nl_hdr: NetlinkHeader,
1461        nl_payload: Vec<u8>,
1462    ) -> Result<NetlinkMessage<GenericMessage>, Errno> {
1463        let Some(status) = audit_status::read_from_bytes(nl_payload.as_bytes()).ok() else {
1464            return error!(EINVAL);
1465        };
1466        self.audit_logger.set_status(current_task, status, self)?;
1467        Ok(AuditNetlinkClient::build_audit_ack(Ok(()), nl_hdr))
1468    }
1469
1470    fn process_user_audit(
1471        &self,
1472        nl_hdr: NetlinkHeader,
1473        nl_payload: Vec<u8>,
1474    ) -> Result<NetlinkMessage<GenericMessage>, Errno> {
1475        let audit_msg = String::from_utf8_lossy(nl_payload.as_bytes());
1476        self.audit_logger.audit_log(nl_hdr.message_type, move || audit_msg);
1477        Ok(AuditNetlinkClient::build_audit_ack(Ok(()), nl_hdr))
1478    }
1479
1480    fn query_events(self: &Arc<Self>) -> FdEvents {
1481        if self.audit_response.lock().is_some() || self.audit_logger.get_backlog_count(self) != 0 {
1482            return FdEvents::POLLIN;
1483        }
1484        FdEvents::empty()
1485    }
1486
1487    fn detach(self: &Arc<Self>) {
1488        self.audit_logger.detach_client(self);
1489    }
1490
1491    fn build_audit_nlmsg(
1492        seq_number: u32,
1493        msg_type: u16,
1494        payload: Vec<u8>,
1495    ) -> NetlinkMessage<GenericMessage> {
1496        // The family in GenericMessage can be used for message type, not only for the Netlink Family,
1497        // because after finalizing the message, the message type is equal to family.
1498        let nl_payload =
1499            NetlinkPayload::InnerMessage(GenericMessage::Other { family: msg_type, payload });
1500        let mut nl_header = NetlinkHeader::default();
1501        nl_header.sequence_number = seq_number;
1502        let mut message = NetlinkMessage::new(nl_header, nl_payload);
1503        message.finalize();
1504        message
1505    }
1506
1507    fn build_audit_ack(
1508        error: Result<(), Errno>,
1509        req_header: NetlinkHeader,
1510    ) -> NetlinkMessage<GenericMessage> {
1511        let error = {
1512            assert_eq!(req_header.buffer_len(), NETLINK_HEADER_LEN);
1513            let mut buffer = vec![0; NETLINK_HEADER_LEN];
1514            req_header.emit(&mut buffer);
1515
1516            let code = match error {
1517                Ok(()) => None,
1518                Err(e) => Some(
1519                    // Audit netlink errors are negative.
1520                    NonZeroI32::new(-(e.code.error_code() as i32))
1521                        .expect("Errno's code must be non-zero"),
1522                ),
1523            };
1524
1525            let mut error = ErrorMessage::default();
1526            error.code = code;
1527            error.header = buffer;
1528            error
1529        };
1530
1531        let payload = NetlinkPayload::<GenericMessage>::Error(error);
1532        let mut resp_header = NetlinkHeader::default();
1533        resp_header.message_type = NLMSG_ERROR;
1534        resp_header.sequence_number = req_header.sequence_number;
1535        let mut message = NetlinkMessage::new(resp_header, payload);
1536        message.finalize();
1537        message
1538    }
1539
1540    fn serialize_nlmsg(message: NetlinkMessage<GenericMessage>) -> Vec<u8> {
1541        let mut buf = vec![0; message.buffer_len()];
1542        message.serialize(&mut buf);
1543        buf
1544    }
1545}
1546
1547/// Audit Netlink Socket structure.
1548pub struct AuditNetlinkSocket {
1549    /// Reference to the `AuditNetlinkClient` associated with self.
1550    audit_client: Arc<AuditNetlinkClient>,
1551}
1552
1553impl AuditNetlinkSocket {
1554    pub fn new(kernel: &Kernel) -> Result<Self, Errno> {
1555        if kernel.audit_logger().is_disabled() {
1556            return error!(EPROTONOSUPPORT);
1557        }
1558        Ok(Self { audit_client: Arc::new(AuditNetlinkClient::new(kernel.audit_logger())) })
1559    }
1560}
1561
1562impl SocketOps for AuditNetlinkSocket {
1563    fn read(
1564        &self,
1565        _socket: &Socket,
1566        _current_task: &CurrentTask,
1567        data: &mut dyn OutputBuffer,
1568        flags: SocketMessageFlags,
1569    ) -> Result<MessageReadInfo, Errno> {
1570        let buf = self.audit_client.read_nlmsg(flags)?;
1571
1572        let size = data.write_all(buf.as_bytes())?;
1573        Ok(MessageReadInfo {
1574            bytes_read: size,
1575            message_length: size,
1576            address: Some(SocketAddress::Netlink(NetlinkAddress::default())),
1577            ancillary_data: vec![],
1578        })
1579    }
1580
1581    fn write(
1582        &self,
1583        socket: &Socket,
1584        current_task: &CurrentTask,
1585        data: &mut dyn InputBuffer,
1586        _dest_address: &mut Option<SocketAddress>,
1587        _ancillary_data: &mut Vec<AncillaryData>,
1588    ) -> Result<usize, Errno> {
1589        match NetlinkMessage::<GenericMessage>::deserialize(
1590            &(data.peek_all()?),
1591            EmptyDeserializeGenlOptions,
1592        ) {
1593            Ok(nl_message) => {
1594                let header = nl_message.header;
1595                security::check_netlink_send_access(current_task, socket, header.message_type)?;
1596
1597                // Send request to the `AuditNetlinkClient`.
1598                let audit_ack = self
1599                    .audit_client
1600                    .process_request(current_task, nl_message)
1601                    .map_err(|e| AuditNetlinkClient::build_audit_ack(Err(e), header))
1602                    .unwrap_or_else(|nlerr| nlerr);
1603                *self.audit_client.audit_response.lock() = Some(audit_ack);
1604                data.drain();
1605                Ok(header.length as usize)
1606            }
1607            Err(e) => {
1608                log_warn!("Failed to process write; data could not be deserialized: {:?}", e);
1609                error!(EINVAL)
1610            }
1611        }
1612    }
1613
1614    fn wait_async(
1615        &self,
1616        _socket: &Socket,
1617        _current_task: &CurrentTask,
1618        waiter: &Waiter,
1619        events: FdEvents,
1620        handler: EventHandler,
1621    ) -> WaitCanceler {
1622        self.audit_client.waiters.wait_async_fd_events(waiter, events, handler)
1623    }
1624
1625    fn query_events(
1626        &self,
1627        _socket: &Socket,
1628        _current_task: &CurrentTask,
1629    ) -> Result<FdEvents, Errno> {
1630        Ok(self.audit_client.query_events() & FdEvents::POLLIN)
1631    }
1632
1633    fn close(&self, _current_task: &CurrentTask, _socket: &Socket) {
1634        // If the `AuditNetlinkClient` disconnects, detach it.
1635        self.audit_client.detach();
1636    }
1637
1638    fn shutdown(&self, _socket: &Socket, _how: SocketShutdownFlags) -> Result<(), Errno> {
1639        error!(EOPNOTSUPP)
1640    }
1641
1642    fn connect(
1643        &self,
1644        _socket: &SocketHandle,
1645        _current_task: &CurrentTask,
1646        _peer: SocketPeer,
1647    ) -> Result<(), Errno> {
1648        error!(EOPNOTSUPP)
1649    }
1650
1651    fn listen(&self, _socket: &Socket, _backlog: i32, _credentials: ucred) -> Result<(), Errno> {
1652        error!(EOPNOTSUPP)
1653    }
1654
1655    fn accept(&self, _socket: &Socket, _current_task: &CurrentTask) -> Result<SocketHandle, Errno> {
1656        error!(EOPNOTSUPP)
1657    }
1658
1659    fn bind(
1660        &self,
1661        _socket: &Socket,
1662        _current_task: &CurrentTask,
1663        _socket_address: SocketAddress,
1664    ) -> Result<(), Errno> {
1665        error!(EOPNOTSUPP)
1666    }
1667
1668    fn getsockname(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
1669        error!(EOPNOTSUPP)
1670    }
1671
1672    fn getpeername(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
1673        error!(EOPNOTSUPP)
1674    }
1675
1676    fn getsockopt(
1677        &self,
1678        _socket: &Socket,
1679        _current_task: &CurrentTask,
1680        _level: u32,
1681        _optname: u32,
1682        _optlen: u32,
1683    ) -> Result<Vec<u8>, Errno> {
1684        error!(EOPNOTSUPP)
1685    }
1686
1687    fn setsockopt(
1688        &self,
1689        _socket: &Socket,
1690        _current_task: &CurrentTask,
1691        _level: u32,
1692        _optname: u32,
1693        _optval: SockOptValue,
1694    ) -> Result<(), Errno> {
1695        error!(EOPNOTSUPP)
1696    }
1697}
1698
1699#[cfg(test)]
1700mod tests {
1701    use super::*;
1702
1703    use netlink_packet_route::route::RouteMessage;
1704    use netlink_packet_route::{RouteNetlinkMessage, RouteNetlinkMessageParseMode};
1705    use test_case::test_case;
1706
1707    // Successfully send the message and observe it's stored in the queue.
1708    #[test_case(true; "sufficient_capacity")]
1709    // Attempting to send when the queue is full should succeed by increasing
1710    // the size of the queue.
1711    #[test_case(false; "insufficient_capacity")]
1712    fn test_netlink_to_client_sender(sufficient_capacity: bool) {
1713        const MODERN_GROUP: u32 = 5;
1714
1715        let mut message: NetlinkMessage<RouteNetlinkMessage> =
1716            RouteNetlinkMessage::NewRoute(RouteMessage::default()).into();
1717        message.finalize();
1718
1719        let (initial_queue_size, final_queue_size) = if sufficient_capacity {
1720            (SOCKET_DEFAULT_SIZE, SOCKET_DEFAULT_SIZE)
1721        } else {
1722            (0, message.buffer_len())
1723        };
1724
1725        let socket_inner = Arc::new(LockDepMutex::new(NetlinkSocketInner {
1726            receive_buffer: MessageQueue::new(initial_queue_size),
1727            ..NetlinkSocketInner::new(NetlinkFamily::Route)
1728        }));
1729
1730        let mut sender = NetlinkToClientSender::<RouteNetlinkMessage>::new(socket_inner.clone());
1731        sender.send(message.clone(), Some(ModernGroup(MODERN_GROUP)));
1732        let Message { data, address, ancillary_data: _ } =
1733            socket_inner.lock().read_message().expect("should read message");
1734
1735        assert_eq!(
1736            address,
1737            Some(SocketAddress::Netlink(NetlinkAddress { pid: 0, groups: 1 << MODERN_GROUP }))
1738        );
1739        let actual_message = NetlinkMessage::<RouteNetlinkMessage>::deserialize(
1740            &data,
1741            RouteNetlinkMessageParseMode::Strict,
1742        )
1743        .expect("message should deserialize into RtnlMessage");
1744        assert_eq!(actual_message, message);
1745        assert_eq!(socket_inner.lock().receive_buffer.capacity(), final_queue_size);
1746    }
1747
1748    fn getsockopt_u32(socket: &NetlinkSocketInner, level: u32, optname: u32) -> u32 {
1749        let byte_vec = socket.getsockopt(level, optname).expect("getsockopt should succeed");
1750        let bytes: [u8; 4] = byte_vec.as_slice().try_into().expect("expected 4 bytes");
1751        u32::from_ne_bytes(bytes)
1752    }
1753
1754    fn sock_opt_value(val: u32) -> SockOptValue {
1755        SockOptValue::Value(val.to_ne_bytes().to_vec())
1756    }
1757
1758    #[::fuchsia::test]
1759    async fn test_set_get_snd_rcv_buf() {
1760        crate::testing::spawn_kernel_and_run_sync(|current_task| {
1761            let mut socket = NetlinkSocketInner::new(NetlinkFamily::Route);
1762
1763            // Verify initialization uses the default value.
1764            let expected_default = u32::try_from(SOCKET_DEFAULT_SIZE).unwrap();
1765            assert_eq!(getsockopt_u32(&socket, SOL_SOCKET, SO_SNDBUF), expected_default);
1766            assert_eq!(getsockopt_u32(&socket, SOL_SOCKET, SO_RCVBUF), expected_default);
1767
1768            // Set new values and observe that they were applied.
1769            // Note that applied value is 2 times the requested value.
1770            const SNDBUF_SIZE: u32 = 12345;
1771            const RCVBUF_SIZE: u32 = 54321;
1772            socket
1773                .setsockopt(current_task, SOL_SOCKET, SO_SNDBUF, sock_opt_value(SNDBUF_SIZE))
1774                .expect("setsockopt should succeed");
1775            socket
1776                .setsockopt(current_task, SOL_SOCKET, SO_RCVBUF, sock_opt_value(RCVBUF_SIZE))
1777                .expect("setsockopt should succeed");
1778            assert_eq!(getsockopt_u32(&socket, SOL_SOCKET, SO_SNDBUF), SNDBUF_SIZE * 2);
1779            assert_eq!(getsockopt_u32(&socket, SOL_SOCKET, SO_RCVBUF), RCVBUF_SIZE * 2);
1780        })
1781        .await;
1782    }
1783
1784    #[::fuchsia::test]
1785    async fn test_snd_rcv_buf_limits() {
1786        crate::testing::spawn_kernel_and_run_sync(|current_task| {
1787            let mut socket = NetlinkSocketInner::new(NetlinkFamily::Route);
1788            let too_big = u32::try_from(SOCKET_MAX_SIZE).unwrap() + 1;
1789
1790            // SO_SNDBUF and SO_RCVBUF clamp the size to the limit.
1791            socket
1792                .setsockopt(current_task, SOL_SOCKET, SO_SNDBUF, sock_opt_value(too_big))
1793                .expect("setsockopt should succeed");
1794            socket
1795                .setsockopt(current_task, SOL_SOCKET, SO_RCVBUF, sock_opt_value(too_big))
1796                .expect("setsockopt should succeed");
1797            let expected_max = u32::try_from(SOCKET_MAX_SIZE).unwrap();
1798            assert_eq!(getsockopt_u32(&socket, SOL_SOCKET, SO_SNDBUF), expected_max);
1799            assert_eq!(getsockopt_u32(&socket, SOL_SOCKET, SO_RCVBUF), expected_max);
1800
1801            // SO_SNDBUFFORCE and SO_RCVBUFFORCE do not.
1802            // Note that the applied value is two times the requested value.
1803            socket
1804                .setsockopt(current_task, SOL_SOCKET, SO_SNDBUFFORCE, sock_opt_value(too_big))
1805                .expect("setsockopt should succeed");
1806            socket
1807                .setsockopt(current_task, SOL_SOCKET, SO_RCVBUFFORCE, sock_opt_value(too_big))
1808                .expect("setsockopt should succeed");
1809            assert_eq!(getsockopt_u32(&socket, SOL_SOCKET, SO_SNDBUF), too_big * 2);
1810            assert_eq!(getsockopt_u32(&socket, SOL_SOCKET, SO_RCVBUF), too_big * 2);
1811        })
1812        .await;
1813    }
1814}