Skip to main content

starnix_core/vfs/socket/
socket_unix.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::bpf::context::EbpfRunContextImpl;
6use crate::bpf::fs::get_bpf_object;
7use crate::mm::MemoryAccessorExt;
8use crate::security;
9use crate::task::{CurrentTask, EventHandler, WaitCanceler, WaitQueue, Waiter};
10use crate::vfs::buffers::{
11    AncillaryData, InputBuffer, MessageQueue, MessageReadInfo, OutputBuffer, UnixControlData,
12};
13use crate::vfs::socket::{
14    AcceptQueue, DEFAULT_LISTEN_BACKLOG, SockOptValue, Socket, SocketAddress, SocketBpfState,
15    SocketDomain, SocketFile, SocketHandle, SocketMessageFlags, SocketOps, SocketPeer,
16    SocketProtocol, SocketShutdownFlags, SocketType,
17};
18use crate::vfs::{
19    AccessCheck, FdNumber, FileHandle, FileObject, FsNodeHandle, FsStr, LookupContext, Message,
20    UcredPtr,
21};
22use ebpf::{
23    BpfProgramContext, BpfValue, CbpfConfig, DataWidth, EbpfProgram, Packet, ProgramArgument, Type,
24};
25use ebpf_api::{
26    LoadBytesBase, PacketWithLoadBytes, PinnedMap, ProgramType, SOCKET_FILTER_CBPF_CONFIG,
27    SOCKET_FILTER_SK_BUF_TYPE, SocketFilterProgramContext, SocketRef,
28};
29use starnix_logging::track_stub;
30use starnix_sync::{LockDepGuard, LockDepMutex, UnixSocketInnerLock, allow_subclass};
31use starnix_syscalls::{SUCCESS, SyscallArg, SyscallResult};
32use starnix_uapi::errors::{EACCES, ECONNREFUSED, EINTR, EPERM, Errno};
33use starnix_uapi::file_mode::Access;
34use starnix_uapi::open_flags::OpenFlags;
35use starnix_uapi::user_address::{UserAddress, UserRef};
36use starnix_uapi::vfs::FdEvents;
37use starnix_uapi::{
38    __sk_buff, FIONREAD, SO_ACCEPTCONN, SO_ATTACH_BPF, SO_BROADCAST, SO_ERROR, SO_KEEPALIVE,
39    SO_LINGER, SO_NO_CHECK, SO_PASSCRED, SO_PASSSEC, SO_PEERCRED, SO_PEERSEC, SO_RCVBUF,
40    SO_REUSEADDR, SO_REUSEPORT, SO_SNDBUF, SOL_SOCKET, errno, error, gid_t, socklen_t, uapi, ucred,
41    uid_t,
42};
43use std::sync::Arc;
44use zerocopy::IntoBytes;
45
46// From unix.go in gVisor.
47const SOCKET_MIN_SIZE: usize = 4 << 10;
48const SOCKET_DEFAULT_SIZE: usize = 208 << 10;
49const SOCKET_MAX_SIZE: usize = 4 << 20;
50
51/// The data of a socket is stored in the "Inner" struct. Because both ends have separate locks,
52/// care must be taken to avoid taking both locks since there is no way to tell what order to
53/// take them in.
54///
55/// When writing, data is buffered in the "other" end of the socket's Inner.MessageQueue:
56///
57///            UnixSocket end #1          UnixSocket end #2
58///            +---------------+          +---------------+
59///            |               |          |   +-------+   |
60///   Writes -------------------------------->| Inner |------> Reads
61///            |               |          |   +-------+   |
62///            |   +-------+   |          |               |
63///   Reads <------| Inner |<-------------------------------- Writes
64///            |   +-------+   |          |               |
65///            +---------------+          +---------------+
66///
67pub struct UnixSocket {
68    inner: LockDepMutex<UnixSocketInner, UnixSocketInnerLock>,
69    waiters: WaitQueue,
70}
71
72fn downcast_socket_to_unix(socket: &Socket) -> &UnixSocket {
73    // It is a programing error if we are downcasting
74    // a different type of socket as sockets from different families
75    // should not communicate, so unwrapping here
76    // will let us know that.
77    socket.downcast_socket::<UnixSocket>().unwrap()
78}
79
80enum UnixSocketState {
81    /// The socket has not been connected.
82    Disconnected,
83
84    /// The socket has had `listen` called and can accept incoming connections.
85    Listening(AcceptQueue),
86
87    /// The socket is connected to a peer.
88    Connected(SocketHandle),
89
90    /// The socket is closed.
91    Closed,
92}
93
94struct UnixSocketInner {
95    /// The `MessageQueue` that contains messages sent to this socket.
96    messages: MessageQueue,
97
98    /// The address that this socket has been bound to, if it has been bound.
99    address: Option<SocketAddress>,
100
101    /// Whether this end of the socket has been shut down and can no longer receive message. It is
102    /// still possible to send messages to the peer, if it exists and hasn't also been shut down.
103    is_read_shutdown: bool,
104
105    /// Whether this end of the socket has been shut down and can no longer send messages.
106    is_write_shutdown: bool,
107
108    /// Whether the peer had unread data when it was closed. In this case, reads should return
109    /// ECONNRESET instead of 0 (eof).
110    peer_closed_with_unread_data: bool,
111
112    /// See SO_LINGER.
113    linger: uapi::linger,
114
115    /// See SO_PASSCRED.
116    passcred: bool,
117
118    /// See SO_PASSSEC.
119    passsec: bool,
120
121    /// See SO_BROADCAST.
122    broadcast: bool,
123
124    /// See SO_NO_CHECK.
125    no_check: bool,
126
127    /// See SO_REUSEPORT.
128    reuseport: bool,
129
130    /// See SO_REUSEADDR.
131    reuseaddr: bool,
132
133    /// See SO_KEEPALIVE.
134    keepalive: bool,
135
136    /// See SO_ATTACH_BPF.
137    bpf_program: Option<UnixSocketFilter>,
138
139    /// Unix credentials of the owner of this socket, for SO_PEERCRED.
140    credentials: Option<ucred>,
141
142    /// Socket state: a queue if this is a listening socket, or a peer if this is a connected
143    /// socket.
144    state: UnixSocketState,
145}
146
147impl UnixSocket {
148    pub fn new(_socket_type: SocketType) -> UnixSocket {
149        UnixSocket {
150            inner: UnixSocketInner {
151                messages: MessageQueue::new(SOCKET_DEFAULT_SIZE),
152                address: None,
153                is_read_shutdown: false,
154                is_write_shutdown: false,
155                peer_closed_with_unread_data: false,
156                linger: uapi::linger::default(),
157                passcred: false,
158                passsec: false,
159                broadcast: false,
160                no_check: false,
161                reuseaddr: false,
162                reuseport: false,
163                keepalive: false,
164                bpf_program: None,
165                credentials: None,
166                state: UnixSocketState::Disconnected,
167            }
168            .into(),
169            waiters: WaitQueue::default(),
170        }
171    }
172
173    /// Creates a pair of connected sockets.
174    ///
175    /// # Parameters
176    /// - `domain`: The domain of the socket (e.g., `AF_UNIX`).
177    /// - `socket_type`: The type of the socket (e.g., `SOCK_STREAM`).
178    pub fn new_pair(
179        current_task: &CurrentTask,
180        domain: SocketDomain,
181        socket_type: SocketType,
182        open_flags: OpenFlags,
183    ) -> Result<(FileHandle, FileHandle), Errno> {
184        let credentials = current_task.current_ucred();
185        let left = Socket::new(
186            current_task,
187            domain,
188            socket_type,
189            SocketProtocol::default(),
190            /* kernel_private = */ false,
191        )?;
192        let right = Socket::new(
193            current_task,
194            domain,
195            socket_type,
196            SocketProtocol::default(),
197            /* kernel_private = */ false,
198        )?;
199        downcast_socket_to_unix(&left).lock().state = UnixSocketState::Connected(right.clone());
200        downcast_socket_to_unix(&left).lock().credentials = Some(credentials.clone());
201        downcast_socket_to_unix(&right).lock().state = UnixSocketState::Connected(left.clone());
202        downcast_socket_to_unix(&right).lock().credentials = Some(credentials);
203        left.set_bpf_state(SocketBpfState::Established);
204        right.set_bpf_state(SocketBpfState::Established);
205        let left = SocketFile::from_socket(
206            current_task,
207            left,
208            open_flags,
209            /* kernel_private= */ false,
210        )?;
211        let right = SocketFile::from_socket(
212            current_task,
213            right,
214            open_flags,
215            /* kernel_private= */ false,
216        )?;
217        let left_socket = SocketFile::get_from_file(&left)?;
218        let right_socket = SocketFile::get_from_file(&right)?;
219
220        security::socket_socketpair(current_task, left_socket, right_socket)?;
221        Ok((left, right))
222    }
223
224    fn connect_stream(
225        &self,
226        socket: &SocketHandle,
227        current_task: &CurrentTask,
228        peer: &SocketHandle,
229    ) -> Result<(), Errno> {
230        // Only hold one lock at a time until we make sure the lock ordering
231        // is right: client before listener
232        match downcast_socket_to_unix(peer).lock().state {
233            UnixSocketState::Listening(_) => {}
234            _ => return error!(ECONNREFUSED),
235        }
236
237        let mut client = downcast_socket_to_unix(socket).lock();
238        match client.state {
239            UnixSocketState::Disconnected => {}
240            UnixSocketState::Connected(_) => return error!(EISCONN),
241            _ => return error!(EINVAL),
242        };
243
244        let unix_socket_peer = downcast_socket_to_unix(peer);
245        {
246            // Lock ordering is client before listener.
247            let _token = allow_subclass();
248            let mut listener = unix_socket_peer.lock();
249
250            // Must check this again because we released the listener lock for a moment
251            let queue = match &listener.state {
252                UnixSocketState::Listening(queue) => queue,
253                _ => return error!(ECONNREFUSED),
254            };
255
256            self.check_type_for_connect(socket, peer, &listener.address)?;
257
258            if queue.sockets.len() > queue.backlog {
259                return error!(EAGAIN);
260            }
261
262            let server = Socket::new(
263                current_task,
264                peer.domain,
265                peer.socket_type,
266                SocketProtocol::default(),
267                /* kernel_private = */ true,
268            )?;
269            security::unix_stream_connect(current_task, socket, peer, &server)?;
270            client.state = UnixSocketState::Connected(server.clone());
271            client.credentials = Some(current_task.current_ucred());
272            {
273                // This allow_subclass is safe because `server` is a newly created socket
274                // that hasn't been added to any public table or returned to the user yet.
275                // It is unreachable by other threads, making lock ordering cycles impossible.
276                let _token = allow_subclass();
277                let mut server = downcast_socket_to_unix(&server).lock();
278                server.state = UnixSocketState::Connected(socket.clone());
279                server.address = listener.address.clone();
280                server.messages.set_capacity(listener.messages.capacity())?;
281                server.credentials = listener.credentials.clone();
282                server.passcred = listener.passcred;
283                server.passsec = listener.passsec;
284            }
285
286            // We already checked that the socket is in Listening state...but the borrow checker cannot
287            // be convinced that it's ok to combine these checks
288            let queue = match listener.state {
289                UnixSocketState::Listening(ref mut queue) => queue,
290                _ => panic!("something changed the server socket state while I held a lock on it"),
291            };
292            queue.sockets.push_back(server);
293        }
294        unix_socket_peer.waiters.notify_fd_events(FdEvents::POLLIN);
295        Ok(())
296    }
297
298    fn connect_datagram(
299        &self,
300        socket: &SocketHandle,
301        current_task: &CurrentTask,
302        peer: &SocketHandle,
303    ) -> Result<(), Errno> {
304        {
305            let unix_socket = socket.downcast_socket::<UnixSocket>().unwrap();
306            let peer_inner = unix_socket.lock();
307            self.check_type_for_connect(socket, peer, &peer_inner.address)?;
308        }
309        security::unix_may_send(current_task, socket, peer)?;
310        let unix_socket = socket.downcast_socket::<UnixSocket>().unwrap();
311        unix_socket.lock().state = UnixSocketState::Connected(peer.clone());
312        Ok(())
313    }
314
315    pub fn check_type_for_connect(
316        &self,
317        socket: &Socket,
318        peer: &Socket,
319        peer_address: &Option<SocketAddress>,
320    ) -> Result<(), Errno> {
321        if socket.domain != peer.domain || socket.socket_type != peer.socket_type {
322            // According to ConnectWithWrongType in accept_bind_test, abstract
323            // UNIX domain sockets return ECONNREFUSED rather than EPROTOTYPE.
324            if let Some(address) = peer_address {
325                if address.is_abstract_unix() {
326                    return error!(ECONNREFUSED);
327                }
328            }
329            return error!(EPROTOTYPE);
330        }
331        Ok(())
332    }
333
334    /// Locks and returns the inner state of the Socket.
335    fn lock(&self) -> LockDepGuard<'_, UnixSocketInner> {
336        self.inner.lock()
337    }
338
339    fn is_listening(&self, _socket: &Socket) -> bool {
340        matches!(self.lock().state, UnixSocketState::Listening(_))
341    }
342
343    fn get_receive_capacity(&self) -> usize {
344        self.lock().messages.capacity()
345    }
346
347    fn set_receive_capacity(&self, requested_capacity: usize) {
348        self.lock().set_capacity(requested_capacity);
349    }
350
351    fn get_send_capacity(&self) -> usize {
352        let peer = {
353            if let Some(peer) = self.lock().peer() {
354                peer.clone()
355            } else {
356                return 0;
357            }
358        };
359        let unix_socket = downcast_socket_to_unix(&peer);
360        let capacity = unix_socket.lock().messages.capacity();
361        capacity
362    }
363
364    fn set_send_capacity(&self, requested_capacity: usize) {
365        let peer = {
366            if let Some(peer) = self.lock().peer() {
367                peer.clone()
368            } else {
369                return;
370            }
371        };
372        let unix_socket = downcast_socket_to_unix(&peer);
373        unix_socket.lock().set_capacity(requested_capacity);
374    }
375
376    fn get_linger(&self) -> uapi::linger {
377        let inner = self.lock();
378        inner.linger
379    }
380
381    fn set_linger(&self, linger: uapi::linger) {
382        let mut inner = self.lock();
383        inner.linger = linger;
384    }
385
386    fn get_passcred(&self) -> bool {
387        let inner = self.lock();
388        inner.passcred
389    }
390
391    fn set_passcred(&self, passcred: bool) {
392        let mut inner = self.lock();
393        inner.passcred = passcred;
394    }
395
396    fn get_passsec(&self) -> bool {
397        let inner = self.lock();
398        inner.passsec
399    }
400
401    fn set_passsec(&self, passsec: bool) {
402        let mut inner = self.lock();
403        inner.passsec = passsec;
404    }
405
406    fn get_broadcast(&self) -> bool {
407        let inner = self.lock();
408        inner.broadcast
409    }
410
411    fn set_broadcast(&self, broadcast: bool) {
412        let mut inner = self.lock();
413        inner.broadcast = broadcast;
414    }
415
416    fn get_no_check(&self) -> bool {
417        let inner = self.lock();
418        inner.no_check
419    }
420
421    fn set_no_check(&self, no_check: bool) {
422        let mut inner = self.lock();
423        inner.no_check = no_check;
424    }
425
426    fn get_reuseaddr(&self) -> bool {
427        let inner = self.lock();
428        inner.reuseaddr
429    }
430
431    fn set_reuseaddr(&self, reuseaddr: bool) {
432        let mut inner = self.lock();
433        inner.reuseaddr = reuseaddr;
434    }
435
436    fn get_reuseport(&self) -> bool {
437        let inner = self.lock();
438        inner.reuseport
439    }
440
441    fn set_reuseport(&self, reuseport: bool) {
442        let mut inner = self.lock();
443        inner.reuseport = reuseport;
444    }
445
446    fn get_keepalive(&self) -> bool {
447        let inner = self.lock();
448        inner.keepalive
449    }
450
451    fn set_keepalive(&self, keepalive: bool) {
452        let mut inner = self.lock();
453        inner.keepalive = keepalive;
454    }
455
456    fn set_bpf_program(&self, program: Option<UnixSocketFilter>) {
457        let mut inner = self.lock();
458        inner.bpf_program = program;
459    }
460
461    fn peer_cred(&self) -> Option<ucred> {
462        let peer = {
463            let inner = self.lock();
464            inner.peer().cloned()
465        };
466        if let Some(peer) = peer {
467            let unix_socket = downcast_socket_to_unix(&peer);
468            let unix_socket = unix_socket.lock();
469            unix_socket.credentials.clone()
470        } else {
471            None
472        }
473    }
474
475    pub fn bind_socket_to_node(
476        &self,
477        socket: &SocketHandle,
478        address: SocketAddress,
479        node: &FsNodeHandle,
480    ) -> Result<(), Errno> {
481        let unix_socket = downcast_socket_to_unix(socket);
482        let mut inner = unix_socket.lock();
483        inner.bind(address)?;
484        node.set_bound_socket(socket.clone());
485        Ok(())
486    }
487
488    fn notify_shutdown(&self) {
489        self.waiters.notify_fd_events(FdEvents::POLLIN | FdEvents::POLLOUT | FdEvents::POLLHUP);
490    }
491}
492
493impl SocketOps for UnixSocket {
494    fn connect(
495        &self,
496        socket: &SocketHandle,
497        current_task: &CurrentTask,
498        peer: SocketPeer,
499    ) -> Result<(), Errno> {
500        let peer = match peer {
501            SocketPeer::Handle(handle) => handle,
502            SocketPeer::Address(SocketAddress::Unspecified) => {
503                if socket.socket_type == SocketType::Datagram {
504                    let unix_socket = socket.downcast_socket::<UnixSocket>().unwrap();
505                    unix_socket.lock().state = UnixSocketState::Disconnected;
506                    return Ok(());
507                }
508                return error!(EINVAL);
509            }
510            SocketPeer::Address(_) => return error!(EINVAL),
511        };
512        match socket.socket_type {
513            SocketType::Stream | SocketType::SeqPacket => {
514                self.connect_stream(socket, current_task, &peer)
515            }
516            SocketType::Datagram | SocketType::Raw => {
517                self.connect_datagram(socket, current_task, &peer)
518            }
519            _ => error!(EINVAL),
520        }
521    }
522
523    fn listen(&self, socket: &Socket, backlog: i32, credentials: ucred) -> Result<(), Errno> {
524        match socket.socket_type {
525            SocketType::Stream | SocketType::SeqPacket => {}
526            _ => return error!(EOPNOTSUPP),
527        }
528        let mut inner = self.lock();
529        inner.credentials = Some(credentials);
530        let is_bound = inner.address.is_some();
531        let backlog = if backlog < 0 { DEFAULT_LISTEN_BACKLOG } else { backlog as usize };
532        match &mut inner.state {
533            UnixSocketState::Disconnected if is_bound => {
534                inner.state = UnixSocketState::Listening(AcceptQueue::new(backlog));
535                Ok(())
536            }
537            UnixSocketState::Listening(queue) => {
538                queue.set_backlog(backlog)?;
539                Ok(())
540            }
541            _ => error!(EINVAL),
542        }
543    }
544
545    fn accept(&self, socket: &Socket, _current_task: &CurrentTask) -> Result<SocketHandle, Errno> {
546        match socket.socket_type {
547            SocketType::Stream | SocketType::SeqPacket => {}
548            _ => return error!(EOPNOTSUPP),
549        }
550        let mut inner = self.lock();
551        let queue = match &mut inner.state {
552            UnixSocketState::Listening(queue) => queue,
553            _ => return error!(EINVAL),
554        };
555        queue.sockets.pop_front().ok_or_else(|| errno!(EAGAIN))
556    }
557
558    fn bind(
559        &self,
560        _socket: &Socket,
561        _current_task: &CurrentTask,
562        socket_address: SocketAddress,
563    ) -> Result<(), Errno> {
564        match socket_address {
565            SocketAddress::Unix(_) => {}
566            _ => return error!(EINVAL),
567        }
568        self.lock().bind(socket_address)
569    }
570
571    fn read(
572        &self,
573        socket: &Socket,
574        _current_task: &CurrentTask,
575        data: &mut dyn OutputBuffer,
576        flags: SocketMessageFlags,
577    ) -> Result<MessageReadInfo, Errno> {
578        let info = self.lock().read(data, socket.socket_type, flags)?;
579        if info.bytes_read > 0 {
580            let peer = {
581                let inner = self.lock();
582                inner.peer().cloned()
583            };
584            if let Some(socket) = peer {
585                let unix_socket_peer = socket.downcast_socket::<UnixSocket>();
586                if let Some(socket) = unix_socket_peer {
587                    socket.waiters.notify_fd_events(FdEvents::POLLOUT);
588                }
589            }
590        }
591        Ok(info)
592    }
593
594    fn write(
595        &self,
596        socket: &Socket,
597        current_task: &CurrentTask,
598        data: &mut dyn InputBuffer,
599        dest_address: &mut Option<SocketAddress>,
600        ancillary_data: &mut Vec<AncillaryData>,
601    ) -> Result<usize, Errno> {
602        let (connected_peer, local_address, creds, is_write_shutdown) = {
603            let inner = self.lock();
604            (
605                inner.peer().map(|p| p.clone()),
606                inner.address.clone(),
607                inner.credentials.clone(),
608                inner.is_write_shutdown,
609            )
610        };
611
612        if is_write_shutdown {
613            return error!(EPIPE);
614        }
615
616        let peer = match (connected_peer, dest_address, socket.socket_type) {
617            (Some(peer), None, _) => peer,
618            (None, Some(_), SocketType::Stream) => return error!(EOPNOTSUPP),
619            (None, Some(_), SocketType::SeqPacket) => return error!(ENOTCONN),
620            (Some(_), Some(_), _) => return error!(EISCONN),
621            (_, Some(SocketAddress::Unix(name)), _) => {
622                resolve_unix_socket_address(current_task, name.as_ref())?
623            }
624            (_, Some(_), _) => return error!(EINVAL),
625            (None, None, _) => return error!(ENOTCONN),
626        };
627
628        if socket.socket_type == SocketType::Datagram {
629            security::unix_may_send(current_task, socket, &peer)?;
630        }
631
632        let unix_socket = downcast_socket_to_unix(&peer);
633        let write_result = {
634            let mut peer = unix_socket.lock();
635            if peer.passcred {
636                let creds = creds.unwrap_or_else(|| current_task.current_ucred());
637                ancillary_data.push(AncillaryData::Unix(UnixControlData::Credentials(creds)));
638            }
639            if socket.socket_type == SocketType::Datagram {
640                // TODO: https://fxbug.dev/364568855 - Store the opaque LSM property value, and expand
641                // it to a string upon readmsg.
642                let context = security::socket_getpeersec_dgram(current_task, socket);
643                ancillary_data.push(AncillaryData::Unix(UnixControlData::Security(context.into())));
644            }
645            peer.write(current_task, data, local_address, ancillary_data, socket.socket_type)
646        };
647
648        if let Err(ref err) = write_result {
649            if err.code == ECONNREFUSED && socket.socket_type == SocketType::Datagram {
650                let mut inner = self.lock();
651                if let UnixSocketState::Connected(ref connected_peer) = inner.state {
652                    if Arc::ptr_eq(connected_peer, &peer) {
653                        inner.state = UnixSocketState::Disconnected;
654                        self.waiters.notify_fd_events(FdEvents::POLLOUT | FdEvents::POLLHUP);
655                    }
656                }
657            }
658        }
659
660        let bytes_written = write_result?;
661        if bytes_written > 0 {
662            unix_socket.waiters.notify_fd_events(FdEvents::POLLIN);
663        }
664        Ok(bytes_written)
665    }
666
667    fn wait_async(
668        &self,
669        _socket: &Socket,
670        _current_task: &CurrentTask,
671        waiter: &Waiter,
672        events: FdEvents,
673        handler: EventHandler,
674    ) -> WaitCanceler {
675        self.waiters.wait_async_fd_events(waiter, events, handler)
676    }
677
678    fn query_events(
679        &self,
680        socket: &Socket,
681        _current_task: &CurrentTask,
682    ) -> Result<FdEvents, Errno> {
683        // Note that self.lock() must be dropped before acquiring peer.inner.lock() to avoid
684        // potential deadlocks.
685        let (
686            mut events,
687            peer,
688            local_is_read_shutdown,
689            local_is_write_shutdown,
690            is_closed,
691            is_disconnected,
692        ) = {
693            let inner = self.lock();
694
695            let mut events = FdEvents::empty();
696            let local_events = inner.messages.query_events();
697            if local_events.contains(FdEvents::POLLIN) {
698                events |= FdEvents::POLLIN;
699            }
700
701            // Listening socket gets POLLIN when there are pending connections.
702            if let UnixSocketState::Listening(queue) = &inner.state {
703                if !queue.sockets.is_empty() {
704                    events |= FdEvents::POLLIN;
705                }
706            }
707
708            (
709                events,
710                inner.peer().cloned(),
711                inner.is_read_shutdown,
712                inner.is_write_shutdown,
713                matches!(inner.state, UnixSocketState::Closed),
714                matches!(inner.state, UnixSocketState::Disconnected),
715            )
716        };
717
718        let connection_oriented = socket.socket_type.is_connection_oriented();
719
720        let mut read_dead = local_is_read_shutdown || is_closed;
721        let mut write_dead = local_is_write_shutdown || is_closed;
722        let mut force_pollout = is_closed || is_disconnected;
723        let mut force_pollhup = false;
724
725        if connection_oriented && is_disconnected {
726            force_pollhup = true;
727            write_dead = true;
728        }
729
730        if let Some(peer) = peer {
731            let unix_socket = downcast_socket_to_unix(&peer);
732            let peer_inner = unix_socket.lock();
733
734            if peer_inner.is_read_shutdown {
735                write_dead = true;
736            }
737            if matches!(peer_inner.state, UnixSocketState::Closed) {
738                read_dead = true;
739                write_dead = true;
740                force_pollout = true;
741            }
742
743            let peer_events = peer_inner.messages.query_events();
744            if peer_events.contains(FdEvents::POLLOUT) {
745                events |= FdEvents::POLLOUT;
746            }
747        }
748
749        if force_pollout {
750            events |= FdEvents::POLLOUT;
751        }
752
753        if read_dead {
754            events |= FdEvents::POLLIN | FdEvents::POLLRDHUP;
755        }
756        if (read_dead && write_dead) || force_pollhup {
757            events |= FdEvents::POLLHUP;
758        }
759
760        Ok(events)
761    }
762
763    /// Shuts down this socket according to how, preventing any future reads and/or writes.
764    ///
765    /// Used by the shutdown syscalls.
766    fn shutdown(&self, socket: &Socket, how: SocketShutdownFlags) -> Result<(), Errno> {
767        let mut self_notify_events = FdEvents::empty();
768        let mut peer_notify_events = FdEvents::empty();
769        let peer = {
770            let mut inner = self.lock();
771            if how.contains(SocketShutdownFlags::READ) {
772                inner.is_read_shutdown = true;
773                self_notify_events |= FdEvents::POLLIN | FdEvents::POLLRDHUP;
774            }
775            if how.contains(SocketShutdownFlags::WRITE) {
776                inner.is_write_shutdown = true;
777                self_notify_events |= FdEvents::POLLOUT;
778            }
779            if inner.is_read_shutdown && inner.is_write_shutdown {
780                self_notify_events |= FdEvents::POLLHUP;
781            }
782            inner.peer().cloned()
783        };
784        if let Some(peer) = &peer {
785            if socket.socket_type.is_connection_oriented() {
786                let unix_socket = downcast_socket_to_unix(peer);
787                let mut peer_inner = unix_socket.lock();
788                if how.contains(SocketShutdownFlags::WRITE) {
789                    peer_inner.is_read_shutdown = true;
790                    peer_notify_events |= FdEvents::POLLIN | FdEvents::POLLRDHUP;
791                }
792                if how.contains(SocketShutdownFlags::READ) {
793                    peer_notify_events |= FdEvents::POLLOUT;
794                }
795                if peer_inner.is_read_shutdown && peer_inner.is_write_shutdown {
796                    peer_notify_events |= FdEvents::POLLHUP;
797                }
798            }
799        }
800        if !self_notify_events.is_empty() {
801            self.waiters.notify_fd_events(self_notify_events);
802        }
803        if !peer_notify_events.is_empty() {
804            if let Some(peer) = &peer {
805                let unix_socket = downcast_socket_to_unix(peer);
806                unix_socket.waiters.notify_fd_events(peer_notify_events);
807            }
808        }
809        Ok(())
810    }
811
812    /// Close this socket.
813    ///
814    /// Called by SocketFile when the file descriptor that is holding this
815    /// socket is closed.
816    ///
817    /// Close differs from shutdown in two ways. First, close will call
818    /// mark_peer_closed_with_unread_data if this socket has unread data,
819    /// which changes how read() behaves on that socket. Second, close
820    /// transitions the internal state of this socket to Closed, which breaks
821    /// the reference cycle that exists in the connected state.
822    fn close(&self, _current_task: &CurrentTask, socket: &Socket) {
823        let (maybe_peer, has_unread) = {
824            let mut inner = self.lock();
825            let maybe_peer = inner.peer().map(Arc::clone);
826            inner.is_read_shutdown = true;
827            inner.state = UnixSocketState::Closed;
828            (maybe_peer, !inner.messages.is_empty())
829        };
830        self.notify_shutdown();
831        // If this is a connected socket type, also shut down the connected peer.
832        if socket.socket_type.is_connection_oriented() {
833            if let Some(peer) = maybe_peer {
834                let unix_socket = downcast_socket_to_unix(&peer);
835
836                {
837                    let mut peer_inner = unix_socket.lock();
838                    if has_unread {
839                        peer_inner.peer_closed_with_unread_data = true;
840                    }
841                    peer_inner.is_read_shutdown = true;
842                }
843                unix_socket.notify_shutdown();
844            }
845        }
846    }
847
848    /// Returns the name of this socket.
849    ///
850    /// The name is derived from the address and domain. A socket
851    /// will always have a name, even if it is not bound to an address.
852    fn getsockname(&self, socket: &Socket) -> Result<SocketAddress, Errno> {
853        let inner = self.lock();
854        if let Some(address) = &inner.address {
855            Ok(address.clone())
856        } else {
857            Ok(SocketAddress::default_for_domain(socket.domain))
858        }
859    }
860
861    /// Returns the name of the peer of this socket, if such a peer exists.
862    ///
863    /// Returns an error if the socket is not connected.
864    fn getpeername(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
865        let peer = self.lock().peer().ok_or_else(|| errno!(ENOTCONN))?.clone();
866        peer.getsockname()
867    }
868
869    fn setsockopt(
870        &self,
871        _socket: &Socket,
872        current_task: &CurrentTask,
873        level: u32,
874        optname: u32,
875        optval: SockOptValue,
876    ) -> Result<(), Errno> {
877        match level {
878            SOL_SOCKET => match optname {
879                SO_SNDBUF => {
880                    let requested_capacity: socklen_t = optval.read(current_task)?;
881                    // See StreamUnixSocketPairTest.SetSocketSendBuf for why we multiply by 2 here.
882                    self.set_send_capacity(requested_capacity as usize * 2);
883                }
884                SO_RCVBUF => {
885                    let requested_capacity: socklen_t = optval.read(current_task)?;
886                    self.set_receive_capacity(requested_capacity as usize);
887                }
888                SO_LINGER => {
889                    let mut linger: uapi::linger = optval.read(current_task)?;
890                    if linger.l_onoff != 0 {
891                        linger.l_onoff = 1;
892                    }
893                    self.set_linger(linger);
894                }
895                SO_PASSCRED => {
896                    let passcred: u32 = optval.read(current_task)?;
897                    self.set_passcred(passcred != 0);
898                }
899                SO_PASSSEC => {
900                    let passsec: u32 = optval.read(current_task)?;
901                    self.set_passsec(passsec != 0);
902                }
903                SO_BROADCAST => {
904                    let broadcast: u32 = optval.read(current_task)?;
905                    self.set_broadcast(broadcast != 0);
906                }
907                SO_NO_CHECK => {
908                    let no_check: u32 = optval.read(current_task)?;
909                    self.set_no_check(no_check != 0);
910                }
911                SO_REUSEADDR => {
912                    let reuseaddr: u32 = optval.read(current_task)?;
913                    self.set_reuseaddr(reuseaddr != 0);
914                }
915                SO_REUSEPORT => {
916                    let reuseport: u32 = optval.read(current_task)?;
917                    self.set_reuseport(reuseport != 0);
918                }
919                SO_KEEPALIVE => {
920                    let keepalive: u32 = optval.read(current_task)?;
921                    self.set_keepalive(keepalive != 0);
922                }
923                SO_ATTACH_BPF => {
924                    let fd: FdNumber = optval.read(current_task)?;
925                    let object = get_bpf_object(current_task, fd)?;
926                    let program = object.as_program()?;
927
928                    let linked_program = program.link(ProgramType::SocketFilter)?;
929
930                    self.set_bpf_program(Some(linked_program));
931                }
932                _ => return error!(ENOPROTOOPT),
933            },
934            _ => return error!(ENOPROTOOPT),
935        }
936        Ok(())
937    }
938
939    fn getsockopt(
940        &self,
941        socket: &Socket,
942        current_task: &CurrentTask,
943        level: u32,
944        optname: u32,
945        _optlen: u32,
946    ) -> Result<Vec<u8>, Errno> {
947        match level {
948            SOL_SOCKET => match optname {
949                SO_PEERCRED => Ok(UcredPtr::into_bytes(
950                    current_task,
951                    self.peer_cred().unwrap_or(ucred { pid: 0, uid: uid_t::MAX, gid: gid_t::MAX }),
952                )
953                .map_err(|_| errno!(EINVAL))?),
954                SO_PEERSEC => match socket.socket_type {
955                    SocketType::Stream => security::socket_getpeersec_stream(current_task, socket),
956                    _ => error!(ENOPROTOOPT),
957                },
958                SO_ACCEPTCONN =>
959                {
960                    #[allow(clippy::bool_to_int_with_if)]
961                    Ok(if self.is_listening(socket) { 1u32 } else { 0u32 }.to_ne_bytes().to_vec())
962                }
963                SO_SNDBUF => Ok((self.get_send_capacity() as socklen_t).to_ne_bytes().to_vec()),
964                SO_RCVBUF => Ok((self.get_receive_capacity() as socklen_t).to_ne_bytes().to_vec()),
965                SO_LINGER => Ok(self.get_linger().as_bytes().to_vec()),
966                SO_PASSCRED => Ok((self.get_passcred() as u32).as_bytes().to_vec()),
967                SO_PASSSEC => Ok((self.get_passsec() as u32).as_bytes().to_vec()),
968                SO_BROADCAST => Ok((self.get_broadcast() as u32).as_bytes().to_vec()),
969                SO_NO_CHECK => Ok((self.get_no_check() as u32).as_bytes().to_vec()),
970                SO_REUSEADDR => Ok((self.get_reuseaddr() as u32).as_bytes().to_vec()),
971                SO_REUSEPORT => Ok((self.get_reuseport() as u32).as_bytes().to_vec()),
972                SO_KEEPALIVE => Ok((self.get_keepalive() as u32).as_bytes().to_vec()),
973                SO_ERROR => Ok((0u32).as_bytes().to_vec()),
974                _ => error!(ENOPROTOOPT),
975            },
976            _ => error!(ENOPROTOOPT),
977        }
978    }
979
980    fn ioctl(
981        &self,
982        socket: &Socket,
983        _file: &FileObject,
984        current_task: &CurrentTask,
985        request: u32,
986        arg: SyscallArg,
987    ) -> Result<SyscallResult, Errno> {
988        let user_addr = UserAddress::from(arg);
989        match request {
990            FIONREAD if socket.socket_type == SocketType::Stream => {
991                let length: i32 =
992                    self.lock().messages.len().try_into().map_err(|_| errno!(EINVAL))?;
993                current_task.write_object(UserRef::<i32>::new(user_addr), &length)?;
994                Ok(SUCCESS)
995            }
996            _ => error!(ENOTTY),
997        }
998    }
999}
1000
1001impl UnixSocketInner {
1002    fn bind(&mut self, socket_address: SocketAddress) -> Result<(), Errno> {
1003        if self.address.is_some() {
1004            return error!(EINVAL);
1005        }
1006        self.address = Some(socket_address);
1007        Ok(())
1008    }
1009
1010    fn set_capacity(&mut self, requested_capacity: usize) {
1011        let capacity = requested_capacity.clamp(SOCKET_MIN_SIZE, SOCKET_MAX_SIZE);
1012        let capacity = std::cmp::max(capacity, self.messages.len());
1013        // We have validated capacity sufficiently that set_capacity should always succeed.
1014        self.messages.set_capacity(capacity).unwrap();
1015    }
1016
1017    /// Returns the socket that is connected to this socket, if such a peer exists. Returns
1018    /// ENOTCONN otherwise.
1019    fn peer(&self) -> Option<&SocketHandle> {
1020        match &self.state {
1021            UnixSocketState::Connected(peer) => Some(peer),
1022            _ => None,
1023        }
1024    }
1025
1026    /// Reads the the contents of this socket into `InputBuffer`.
1027    ///
1028    /// Will stop reading if a message with ancillary data is encountered (after the message with
1029    /// ancillary data has been read).
1030    ///
1031    /// # Parameters
1032    /// - `data`: The `OutputBuffer` to write the data to.
1033    ///
1034    /// Returns the number of bytes that were read into the buffer, and any ancillary data that was
1035    /// read from the socket.
1036    fn read(
1037        &mut self,
1038        data: &mut dyn OutputBuffer,
1039        socket_type: SocketType,
1040        flags: SocketMessageFlags,
1041    ) -> Result<MessageReadInfo, Errno> {
1042        let (mut info, has_message) = if socket_type == SocketType::Stream {
1043            if flags.contains(SocketMessageFlags::PEEK) {
1044                self.messages.peek_stream(data)?
1045            } else {
1046                self.messages.read_stream(data)?
1047            }
1048        } else if flags.contains(SocketMessageFlags::PEEK) {
1049            self.messages.peek_datagram(data)?
1050        } else {
1051            self.messages.read_datagram(data)?
1052        };
1053        if !has_message {
1054            if self.peer_closed_with_unread_data {
1055                // Reset the flag
1056                self.peer_closed_with_unread_data = false;
1057                return error!(ECONNRESET);
1058            }
1059            if !self.is_read_shutdown {
1060                return error!(EAGAIN);
1061            }
1062        }
1063
1064        // Remove any credentials message, so that it can be moved to the front if passcred is
1065        // enabled, or simply be removed if passcred is not enabled.
1066        let creds_message;
1067        if let Some(index) = info
1068            .ancillary_data
1069            .iter()
1070            .position(|m| matches!(m, AncillaryData::Unix(UnixControlData::Credentials { .. })))
1071        {
1072            creds_message = info.ancillary_data.remove(index)
1073        } else {
1074            // If passcred is enabled credentials are returned even if they were not sent.
1075            creds_message = AncillaryData::Unix(UnixControlData::unknown_creds());
1076        }
1077        if self.passcred {
1078            // Allow credentials to take priority if they are enabled, so insert at 0.
1079            info.ancillary_data.insert(0, creds_message);
1080        }
1081
1082        // Security labels are only delivered if passsec is currently enabled on this socket.
1083        if !self.passsec {
1084            info.ancillary_data
1085                .retain(|m| !matches!(m, AncillaryData::Unix(UnixControlData::Security(..))));
1086        }
1087
1088        Ok(info)
1089    }
1090
1091    /// Writes the the contents of `InputBuffer` into this socket.
1092    ///
1093    /// # Parameters
1094    /// - `data`: The `InputBuffer` to read the data from.
1095    /// - `ancillary_data`: Any ancillary data to write to the socket. Note that the ancillary data
1096    ///                     will only be written if the entirety of the requested write completes.
1097    ///
1098    /// Returns the number of bytes that were written to the socket.
1099    fn write(
1100        &mut self,
1101        current_task: &CurrentTask,
1102        data: &mut dyn InputBuffer,
1103        address: Option<SocketAddress>,
1104        ancillary_data: &mut Vec<AncillaryData>,
1105        socket_type: SocketType,
1106    ) -> Result<usize, Errno> {
1107        if matches!(self.state, UnixSocketState::Closed) {
1108            if !socket_type.is_connection_oriented() {
1109                return error!(ECONNREFUSED);
1110            } else {
1111                return error!(EPIPE);
1112            }
1113        }
1114        if self.is_read_shutdown {
1115            return error!(EPIPE);
1116        }
1117        let filter = |mut message: Message| {
1118            let Some(bpf_program) = self.bpf_program.as_ref() else {
1119                return Some(message);
1120            };
1121
1122            // TODO(https://fxbug.dev/385015056): Fill in SkBuf.
1123            let mut sk_buf = SkBuf::default();
1124
1125            let mut context = EbpfRunContextImpl::<'_>::new(current_task);
1126            let s = bpf_program.run(&mut context, &mut sk_buf);
1127            if s == 0 {
1128                None
1129            } else {
1130                message.truncate(s as usize);
1131                Some(message)
1132            }
1133        };
1134        let bytes_written = if socket_type == SocketType::Stream {
1135            self.messages.write_stream_with_filter(data, address, ancillary_data, filter)?
1136        } else {
1137            self.messages.write_datagram_with_filter(data, address, ancillary_data, filter)?
1138        };
1139        Ok(bytes_written)
1140    }
1141}
1142
1143pub fn resolve_unix_socket_address(
1144    current_task: &CurrentTask,
1145    name: &FsStr,
1146) -> Result<SocketHandle, Errno> {
1147    if name[0] == b'\0' {
1148        current_task.running_state().abstract_socket_namespace.lookup(name)
1149    } else {
1150        let mut context = LookupContext::default();
1151        let (parent, basename) =
1152            current_task.lookup_parent_at(&mut context, FdNumber::AT_FDCWD, name)?;
1153        let name = parent.lookup_child(current_task, &mut context, basename).map_err(|errno| {
1154            if matches!(errno.code, EACCES | EPERM | EINTR) { errno } else { errno!(ECONNREFUSED) }
1155        })?;
1156        name.check_access(current_task, AccessCheck::for_internal(Access::WRITE))?;
1157        name.entry.node.bound_socket().map(|s| s.clone()).ok_or_else(|| errno!(ECONNREFUSED))
1158    }
1159}
1160
1161// Packet buffer representation used for eBPF filters.
1162#[repr(C)]
1163#[derive(Default)]
1164struct SkBuf {
1165    sk_buff: __sk_buff,
1166}
1167
1168impl Packet for &mut SkBuf {
1169    fn load(&self, _offset: i32, _width: DataWidth) -> Option<BpfValue> {
1170        // TODO(https://fxbug.dev/385015056): Implement packet access.
1171        None
1172    }
1173}
1174
1175impl<'a> PacketWithLoadBytes for &'a mut SkBuf {
1176    fn load_bytes_relative(
1177        &self,
1178        _base: LoadBytesBase,
1179        _offset: usize,
1180        _buf: ebpf::EbpfBufferPtr<'_>,
1181    ) -> i64 {
1182        track_stub!(TODO("https://fxbug.dev/385015056"), "bpf_load_bytes_relative");
1183        -1
1184    }
1185}
1186
1187impl ProgramArgument for &'_ mut SkBuf {
1188    fn get_type() -> &'static Type {
1189        &*SOCKET_FILTER_SK_BUF_TYPE
1190    }
1191}
1192
1193impl SocketRef for &'_ mut SkBuf {
1194    fn get_socket_cookie(&self) -> Option<u64> {
1195        track_stub!(TODO("https://fxbug.dev/385015056"), "bpf_get_socket_cookie");
1196        None
1197    }
1198
1199    fn get_socket_uid(&self) -> Option<uid_t> {
1200        track_stub!(TODO("https://fxbug.dev/385015056"), "bpf_get_socket_uid");
1201        None
1202    }
1203}
1204
1205struct UnixSocketEbpfContext {}
1206impl BpfProgramContext for UnixSocketEbpfContext {
1207    type RunContext<'a> = EbpfRunContextImpl<'a>;
1208    type Packet<'a> = &'a mut SkBuf;
1209    type Map = PinnedMap;
1210    const CBPF_CONFIG: &'static CbpfConfig = &SOCKET_FILTER_CBPF_CONFIG;
1211}
1212
1213ebpf_api::ebpf_program_context_type!(UnixSocketEbpfContext, SocketFilterProgramContext);
1214
1215type UnixSocketFilter = EbpfProgram<UnixSocketEbpfContext>;
1216
1217#[cfg(test)]
1218mod tests {
1219    use super::*;
1220    use crate::mm::MemoryAccessor;
1221    use crate::testing::{map_memory, spawn_kernel_and_run};
1222    use starnix_types::user_buffer::UserBuffer;
1223
1224    #[::fuchsia::test]
1225    async fn test_socket_send_capacity() {
1226        spawn_kernel_and_run(async |current_task| {
1227            let _kernel = current_task.kernel();
1228            let socket = Socket::new(
1229                &current_task,
1230                SocketDomain::Unix,
1231                SocketType::Stream,
1232                SocketProtocol::default(),
1233                /* kernel_private = */ false,
1234            )
1235            .expect("Failed to create socket.");
1236            socket
1237                .bind(&current_task, SocketAddress::Unix(b"\0".into()))
1238                .expect("Failed to bind socket.");
1239            socket.listen(&current_task, 10).expect("Failed to listen.");
1240            let connecting_socket = Socket::new(
1241                &current_task,
1242                SocketDomain::Unix,
1243                SocketType::Stream,
1244                SocketProtocol::default(),
1245                /* kernel_private = */ false,
1246            )
1247            .expect("Failed to connect socket.");
1248            connecting_socket
1249                .ops
1250                .connect(&connecting_socket, &current_task, SocketPeer::Handle(socket.clone()))
1251                .expect("Failed to connect socket.");
1252            assert_eq!(Ok(FdEvents::POLLIN), socket.query_events(&current_task));
1253            let server_socket = socket.accept(&current_task).unwrap();
1254
1255            let opt_size = std::mem::size_of::<socklen_t>();
1256            let user_address = map_memory(&current_task, UserAddress::default(), opt_size as u64);
1257            let send_capacity: socklen_t = 4 * 4096;
1258            current_task.write_memory(user_address, &send_capacity.to_ne_bytes()).unwrap();
1259            let user_buffer = UserBuffer { address: user_address, length: opt_size };
1260            server_socket
1261                .setsockopt(&current_task, SOL_SOCKET, SO_SNDBUF, user_buffer.into())
1262                .unwrap();
1263
1264            let opt_bytes =
1265                server_socket.getsockopt(&current_task, SOL_SOCKET, SO_SNDBUF, 0).unwrap();
1266            let retrieved_capacity = socklen_t::from_ne_bytes(opt_bytes.try_into().unwrap());
1267            // Setting SO_SNDBUF actually sets it to double the size
1268            assert_eq!(2 * send_capacity, retrieved_capacity);
1269        })
1270        .await;
1271    }
1272
1273    #[::fuchsia::test]
1274    async fn test_datagram_socket_disconnect_af_unspec() {
1275        spawn_kernel_and_run(async |current_task| {
1276            let socket1 = Socket::new(
1277                &current_task,
1278                SocketDomain::Unix,
1279                SocketType::Datagram,
1280                SocketProtocol::default(),
1281                /* kernel_private = */ false,
1282            )
1283            .expect("Failed to create socket 1.");
1284            socket1
1285                .bind(&current_task, SocketAddress::Unix(b"\0sock1".into()))
1286                .expect("Failed to bind socket 1.");
1287
1288            let socket2 = Socket::new(
1289                &current_task,
1290                SocketDomain::Unix,
1291                SocketType::Datagram,
1292                SocketProtocol::default(),
1293                /* kernel_private = */ false,
1294            )
1295            .expect("Failed to create socket 2.");
1296            socket2
1297                .bind(&current_task, SocketAddress::Unix(b"\0sock2".into()))
1298                .expect("Failed to bind socket 2.");
1299
1300            // Connect socket1 to socket2.
1301            socket1
1302                .ops
1303                .connect(&socket1, &current_task, SocketPeer::Handle(socket2.clone()))
1304                .expect("Failed to connect socket1 to socket2.");
1305
1306            // Disconnect socket1 using AF_UNSPEC.
1307            socket1
1308                .ops
1309                .connect(&socket1, &current_task, SocketPeer::Address(SocketAddress::Unspecified))
1310                .expect("Failed to disconnect socket1 with AF_UNSPEC.");
1311        })
1312        .await;
1313    }
1314}