Skip to main content

starnix_core/vfs/socket/
socket_vsock.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::task::{CurrentTask, EventHandler, WaitCanceler, WaitQueue, Waiter};
6use crate::vfs::FileHandle;
7use crate::vfs::buffers::{AncillaryData, InputBuffer, MessageReadInfo, OutputBuffer};
8use crate::vfs::socket::{
9    AcceptQueue, DEFAULT_LISTEN_BACKLOG, Socket, SocketAddress, SocketDomain, SocketHandle,
10    SocketMessageFlags, SocketOps, SocketPeer, SocketProtocol, SocketShutdownFlags, SocketType,
11};
12use starnix_sync::{LockDepGuard, LockDepMutex, VsockSocketInnerLock, allow_subclass};
13use starnix_uapi::auth::Credentials;
14use starnix_uapi::errors::Errno;
15use starnix_uapi::open_flags::OpenFlags;
16use starnix_uapi::vfs::FdEvents;
17use starnix_uapi::{errno, error, ucred};
18
19// An implementation of AF_VSOCK.
20// See https://man7.org/linux/man-pages/man7/vsock.7.html
21
22pub struct VsockSocket {
23    inner: LockDepMutex<VsockSocketInner, VsockSocketInnerLock>,
24}
25
26struct VsockSocketInner {
27    /// The address that this socket has been bound to, if it has been bound.
28    address: Option<SocketAddress>,
29
30    // WaitQueue for listening sockets.
31    waiters: WaitQueue,
32
33    // state of the vsock. Contains a handle to a ZxioBackedSocket when connected.
34    state: VsockSocketState,
35}
36
37enum VsockSocketState {
38    /// The socket has not been connected.
39    Disconnected,
40
41    /// The socket has had `listen` called and can accept incoming connections.
42    Listening(AcceptQueue),
43
44    /// The socket is connected to a ZxioBackedSocket.
45    Connected { file: FileHandle, peer_addr: SocketAddress },
46
47    /// The socket is closed.
48    Closed,
49}
50
51fn downcast_socket_to_vsock(socket: &Socket) -> &VsockSocket {
52    // It is a programing error if we are downcasting
53    // a different type of socket as sockets from different families
54    // should not communicate, so unwrapping here
55    // will let us know that.
56    socket.downcast_socket::<VsockSocket>().unwrap()
57}
58
59impl VsockSocket {
60    pub fn new(_socket_type: SocketType) -> VsockSocket {
61        VsockSocket {
62            inner: VsockSocketInner {
63                address: None,
64                waiters: WaitQueue::default(),
65                state: VsockSocketState::Disconnected,
66            }
67            .into(),
68        }
69    }
70
71    /// Locks and returns the inner state of the Socket.
72    fn lock(&self) -> LockDepGuard<'_, VsockSocketInner> {
73        self.inner.lock()
74    }
75}
76
77impl SocketOps for VsockSocket {
78    // Connect with Vsock sockets is not allowed as
79    // we only connect from the enclosing OK.
80    fn connect(
81        &self,
82        _socket: &SocketHandle,
83        _current_task: &CurrentTask,
84        _peer: SocketPeer,
85    ) -> Result<(), Errno> {
86        error!(EPROTOTYPE)
87    }
88
89    fn listen(&self, _socket: &Socket, backlog: i32, _credentials: ucred) -> Result<(), Errno> {
90        let mut inner = self.lock();
91        let is_bound = inner.address.is_some();
92        let backlog = if backlog < 0 { DEFAULT_LISTEN_BACKLOG } else { backlog as usize };
93        match &mut inner.state {
94            VsockSocketState::Disconnected if is_bound => {
95                inner.state = VsockSocketState::Listening(AcceptQueue::new(backlog));
96                Ok(())
97            }
98            VsockSocketState::Listening(queue) => {
99                queue.set_backlog(backlog)?;
100                Ok(())
101            }
102            _ => error!(EINVAL),
103        }
104    }
105
106    fn accept(&self, socket: &Socket, _current_task: &CurrentTask) -> Result<SocketHandle, Errno> {
107        match socket.socket_type {
108            SocketType::Stream | SocketType::SeqPacket => {}
109            _ => return error!(EOPNOTSUPP),
110        }
111        let mut inner = self.lock();
112        let queue = match &mut inner.state {
113            VsockSocketState::Listening(queue) => queue,
114            _ => return error!(EINVAL),
115        };
116        let socket = queue.sockets.pop_front().ok_or_else(|| errno!(EAGAIN))?;
117        Ok(socket)
118    }
119
120    fn bind(
121        &self,
122        _socket: &Socket,
123        _current_task: &CurrentTask,
124        socket_address: SocketAddress,
125    ) -> Result<(), Errno> {
126        match socket_address {
127            SocketAddress::Vsock { .. } => {}
128            _ => return error!(EINVAL),
129        }
130        let mut inner = self.lock();
131        if inner.address.is_some() {
132            return error!(EINVAL);
133        }
134        inner.address = Some(socket_address);
135        Ok(())
136    }
137
138    fn read(
139        &self,
140        _socket: &Socket,
141        current_task: &CurrentTask,
142        data: &mut dyn OutputBuffer,
143        _flags: SocketMessageFlags,
144    ) -> Result<MessageReadInfo, Errno> {
145        let (address, file) = {
146            let inner = self.lock();
147            let address = inner.address.clone();
148
149            match &inner.state {
150                VsockSocketState::Connected { file, .. } => (address, file.clone()),
151                _ => return error!(EBADF),
152            }
153        };
154        let bytes_read =
155            current_task.override_creds(Credentials::root(), || file.read(current_task, data))?;
156        Ok(MessageReadInfo {
157            bytes_read,
158            message_length: bytes_read,
159            address,
160            ancillary_data: vec![],
161        })
162    }
163
164    fn write(
165        &self,
166        _socket: &Socket,
167        current_task: &CurrentTask,
168        data: &mut dyn InputBuffer,
169        _dest_address: &mut Option<SocketAddress>,
170        _ancillary_data: &mut Vec<AncillaryData>,
171    ) -> Result<usize, Errno> {
172        let file = {
173            let inner = self.lock();
174            match &inner.state {
175                VsockSocketState::Connected { file, .. } => file.clone(),
176                _ => return error!(EBADF),
177            }
178        };
179        current_task.override_creds(Credentials::root(), || file.write(current_task, data))
180    }
181
182    fn wait_async(
183        &self,
184        _socket: &Socket,
185        current_task: &CurrentTask,
186        waiter: &Waiter,
187        events: FdEvents,
188        handler: EventHandler,
189    ) -> WaitCanceler {
190        let inner = self.lock();
191        match &inner.state {
192            VsockSocketState::Connected { file, .. } => file
193                .wait_async(current_task, waiter, events, handler)
194                .expect("vsock socket should be connected to a file that can be waited on"),
195            _ => inner.waiters.wait_async_fd_events(waiter, events, handler),
196        }
197    }
198
199    fn query_events(
200        &self,
201        _socket: &Socket,
202        current_task: &CurrentTask,
203    ) -> Result<FdEvents, Errno> {
204        self.lock().query_events(current_task)
205    }
206
207    fn shutdown(&self, _socket: &Socket, _how: SocketShutdownFlags) -> Result<(), Errno> {
208        self.lock().state = VsockSocketState::Closed;
209        Ok(())
210    }
211
212    fn close(&self, _current_task: &CurrentTask, socket: &Socket) {
213        // Call to shutdown should never fail, so unwrap is OK
214        self.shutdown(socket, SocketShutdownFlags::READ | SocketShutdownFlags::WRITE).unwrap();
215    }
216
217    fn getsockname(&self, socket: &Socket) -> Result<SocketAddress, Errno> {
218        let inner = self.lock();
219        if let Some(address) = &inner.address {
220            Ok(address.clone())
221        } else {
222            Ok(SocketAddress::default_for_domain(socket.domain))
223        }
224    }
225
226    fn getpeername(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
227        let inner = self.lock();
228        match &inner.state {
229            VsockSocketState::Connected { peer_addr, .. } => Ok(peer_addr.clone()),
230            _ => {
231                error!(ENOTCONN)
232            }
233        }
234    }
235}
236
237impl VsockSocket {
238    pub fn remote_connection(
239        &self,
240        socket: &Socket,
241        current_task: &CurrentTask,
242        file: FileHandle,
243    ) -> Result<(), Errno> {
244        // we only allow non-blocking files here, so that
245        // read and write on file can return EAGAIN.
246        assert!(file.flags().contains(OpenFlags::NONBLOCK));
247        if socket.socket_type != SocketType::Stream {
248            return error!(ENOTSUP);
249        }
250        if socket.domain != SocketDomain::Vsock {
251            return error!(EINVAL);
252        }
253
254        let mut inner = self.lock();
255        match &mut inner.state {
256            VsockSocketState::Listening(queue) => {
257                if queue.sockets.len() >= queue.backlog {
258                    return error!(EAGAIN);
259                }
260                let remote_socket = Socket::new(
261                    current_task,
262                    SocketDomain::Vsock,
263                    SocketType::Stream,
264                    SocketProtocol::default(),
265                    /* kernel_private = */ true,
266                )?;
267                // The lock on remote_socket is safe because the socket is not shared with any other
268                // thread yet.
269                let _token = allow_subclass();
270                downcast_socket_to_vsock(&remote_socket).lock().state =
271                    VsockSocketState::Connected {
272                        file,
273                        peer_addr: SocketAddress::Vsock {
274                            port: u32::MAX,
275                            cid: starnix_uapi::VMADDR_CID_HOST,
276                        },
277                    };
278                queue.sockets.push_back(remote_socket);
279                inner.waiters.notify_fd_events(FdEvents::POLLIN);
280                Ok(())
281            }
282            _ => error!(EINVAL),
283        }
284    }
285}
286
287impl VsockSocketInner {
288    fn query_events(&self, current_task: &CurrentTask) -> Result<FdEvents, Errno> {
289        Ok(match &self.state {
290            VsockSocketState::Disconnected => FdEvents::empty(),
291            VsockSocketState::Connected { file, .. } => current_task
292                .override_creds(Credentials::root(), || file.query_events(current_task))?,
293            VsockSocketState::Listening(queue) => {
294                if !queue.sockets.is_empty() {
295                    FdEvents::POLLIN
296                } else {
297                    FdEvents::empty()
298                }
299            }
300            VsockSocketState::Closed => FdEvents::POLLHUP,
301        })
302    }
303}
304
305#[cfg(test)]
306mod tests {
307    use super::*;
308    use crate::fs::fuchsia::create_fuchsia_pipe;
309    use crate::mm::PAGE_SIZE;
310    use crate::task::dynamic_thread_spawner::SpawnRequestBuilder;
311    use crate::testing::spawn_kernel_and_run;
312    use crate::vfs::EpollFileObject;
313    use crate::vfs::buffers::{VecInputBuffer, VecOutputBuffer};
314    use crate::vfs::socket::SocketFile;
315    use futures::executor::block_on;
316    use starnix_uapi::vfs::EpollEvent;
317    use syncio::Zxio;
318
319    #[::fuchsia::test]
320    async fn test_vsock_socket() {
321        spawn_kernel_and_run(async |current_task| {
322            let (fs1, fs2) = fidl::Socket::create_stream();
323            const VSOCK_PORT: u32 = 5555;
324
325            let listen_socket = Socket::new(
326                &current_task,
327                SocketDomain::Vsock,
328                SocketType::Stream,
329                SocketProtocol::default(),
330                /* kernel_private = */ false,
331            )
332            .expect("Failed to create socket.");
333            let vsock_ns = current_task.running_state().abstract_vsock_namespace.clone();
334            vsock_ns
335                .bind(&current_task, VSOCK_PORT, &listen_socket)
336                .expect("Failed to bind socket.");
337            listen_socket.listen(&current_task, 10).expect("Failed to listen.");
338
339            let listen_socket =
340                vsock_ns.lookup(&VSOCK_PORT).expect("Failed to look up listening socket.");
341            let remote =
342                create_fuchsia_pipe(&current_task, fs2, OpenFlags::RDWR | OpenFlags::NONBLOCK)
343                    .unwrap();
344            listen_socket
345                .downcast_socket::<VsockSocket>()
346                .unwrap()
347                .remote_connection(&listen_socket, &current_task, remote)
348                .unwrap();
349
350            let server_socket = listen_socket.accept(&current_task).unwrap();
351
352            let test_bytes_in: [u8; 5] = [0, 1, 2, 3, 4];
353            assert_eq!(fs1.write(&test_bytes_in[..]).unwrap(), test_bytes_in.len());
354            let mut buffer_iterator = VecOutputBuffer::new(*PAGE_SIZE as usize);
355            let read_message_info = server_socket
356                .read(&current_task, &mut buffer_iterator, SocketMessageFlags::empty())
357                .unwrap();
358            assert_eq!(read_message_info.bytes_read, test_bytes_in.len());
359            assert_eq!(buffer_iterator.data(), test_bytes_in);
360
361            let test_bytes_out: [u8; 10] = [9, 8, 7, 6, 5, 4, 3, 2, 1, 0];
362            let mut buffer_iterator = VecInputBuffer::new(&test_bytes_out);
363            server_socket
364                .write(&current_task, &mut buffer_iterator, &mut None, &mut vec![])
365                .unwrap();
366            assert_eq!(buffer_iterator.bytes_read(), test_bytes_out.len());
367
368            let mut read_back_buf = [0u8; 100];
369            assert_eq!(test_bytes_out.len(), fs1.read(&mut read_back_buf).unwrap());
370            assert_eq!(&read_back_buf[..test_bytes_out.len()], &test_bytes_out);
371
372            server_socket.close(&current_task);
373            listen_socket.close(&current_task);
374        })
375        .await;
376    }
377
378    #[::fuchsia::test]
379    async fn test_vsock_write_while_read() {
380        spawn_kernel_and_run(async |current_task| {
381            let kernel = current_task.kernel();
382            let (fs1, fs2) = fidl::Socket::create_stream();
383            let socket = Socket::new(
384                &current_task,
385                SocketDomain::Vsock,
386                SocketType::Stream,
387                SocketProtocol::default(),
388                /* kernel_private = */ false,
389            )
390            .expect("Failed to create socket.");
391            let remote =
392                create_fuchsia_pipe(&current_task, fs2, OpenFlags::RDWR | OpenFlags::NONBLOCK)
393                    .unwrap();
394            downcast_socket_to_vsock(&socket).lock().state = VsockSocketState::Connected {
395                file: remote,
396                peer_addr: SocketAddress::Vsock {
397                    port: u32::MAX,
398                    cid: starnix_uapi::VMADDR_CID_HOST,
399                },
400            };
401            let socket_file =
402                SocketFile::from_socket(&current_task, socket, OpenFlags::RDWR, false)
403                    .expect("Failed to create socket file.");
404
405            const XFER_SIZE: usize = 42;
406
407            let socket_clone = socket_file.clone();
408            let closure = move |current_task: &CurrentTask| {
409                let bytes_read =
410                    socket_clone.read(current_task, &mut VecOutputBuffer::new(XFER_SIZE)).unwrap();
411                assert_eq!(XFER_SIZE, bytes_read);
412            };
413            let (result, req) =
414                SpawnRequestBuilder::new().with_sync_closure(closure).build_with_async_result();
415            kernel.kthreads.spawner().spawn_from_request(req);
416
417            // Wait for the thread to become blocked on the read.
418            std::thread::sleep(std::time::Duration::from_secs(2));
419
420            socket_file.write(&current_task, &mut VecInputBuffer::new(&[0; XFER_SIZE])).unwrap();
421
422            let mut buffer = [0u8; 1024];
423            assert_eq!(XFER_SIZE, fs1.read(&mut buffer).unwrap());
424            assert_eq!(XFER_SIZE, fs1.write(&buffer[..XFER_SIZE]).unwrap());
425            block_on(result).unwrap();
426        })
427        .await;
428    }
429
430    #[::fuchsia::test]
431    async fn test_vsock_poll() {
432        spawn_kernel_and_run(async |current_task| {
433            let (client, server) = zx::Socket::create_stream();
434            let pipe = create_fuchsia_pipe(&current_task, client, OpenFlags::RDWR)
435                .expect("create_fuchsia_pipe");
436            let server_zxio = Zxio::create(server.into_handle()).expect("Zxio::create");
437            let socket_object = Socket::new(
438                &current_task,
439                SocketDomain::Vsock,
440                SocketType::Stream,
441                SocketProtocol::default(),
442                /* kernel_private = */ false,
443            )
444            .expect("Failed to create socket.");
445            downcast_socket_to_vsock(&socket_object).lock().state = VsockSocketState::Connected {
446                file: pipe,
447                peer_addr: SocketAddress::Vsock {
448                    port: u32::MAX,
449                    cid: starnix_uapi::VMADDR_CID_HOST,
450                },
451            };
452            let socket = SocketFile::from_socket(
453                &current_task,
454                socket_object.clone(),
455                OpenFlags::RDWR,
456                false,
457            )
458            .expect("Failed to create socket file.");
459
460            assert_eq!(
461                socket.query_events(&current_task),
462                Ok(FdEvents::POLLOUT | FdEvents::POLLWRNORM)
463            );
464
465            let epoll_object = EpollFileObject::new_file(&current_task);
466            let epoll_file = epoll_object.downcast_file::<EpollFileObject>().unwrap();
467            let event = EpollEvent::new(FdEvents::POLLIN, 0);
468            epoll_file.add(&current_task, &socket, &epoll_object, event).expect("poll_file.add");
469
470            let fds = epoll_file.wait(&current_task, 1, zx::MonotonicInstant::ZERO).expect("wait");
471            assert!(fds.is_empty());
472
473            assert_eq!(server_zxio.write(&[0]).expect("write"), 1);
474
475            assert_eq!(
476                socket.query_events(&current_task),
477                Ok(FdEvents::POLLOUT
478                    | FdEvents::POLLWRNORM
479                    | FdEvents::POLLIN
480                    | FdEvents::POLLRDNORM)
481            );
482            let fds = epoll_file.wait(&current_task, 1, zx::MonotonicInstant::ZERO).expect("wait");
483            assert_eq!(fds.len(), 1);
484
485            assert_eq!(socket.read(&current_task, &mut VecOutputBuffer::new(64)).expect("read"), 1);
486
487            assert_eq!(
488                socket.query_events(&current_task),
489                Ok(FdEvents::POLLOUT | FdEvents::POLLWRNORM)
490            );
491            let fds = epoll_file.wait(&current_task, 1, zx::MonotonicInstant::ZERO).expect("wait");
492            assert!(fds.is_empty());
493        })
494        .await;
495    }
496}