Skip to main content

starnix_core/fs/fuchsia/
remote_unix_domain_socket.rs

1// Copyright 2024 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::fs::fuchsia::{OpenFlags, new_remote_file};
6use crate::task::{
7    CurrentTask, EventHandler, SignalHandler, SignalHandlerInner, WaitCanceler, Waiter,
8};
9use crate::vfs::buffers::{InputBuffer, OutputBuffer};
10use crate::vfs::socket::{
11    SockOptValue, Socket, SocketAddress, SocketDomain, SocketHandle, SocketMessageFlags, SocketOps,
12    SocketPeer, SocketProtocol, SocketShutdownFlags, SocketType,
13};
14use crate::vfs::{AncillaryData, FileHandle, MessageReadInfo, UnixControlData};
15use fidl::endpoints::SynchronousProxy;
16use fidl_fuchsia_io as fio;
17use fidl_fuchsia_starnix_binder as fbinder;
18use linux_uapi::{SO_LINGER, SOL_SOCKET};
19use starnix_uapi::auth::Credentials;
20use starnix_uapi::errors::Errno;
21use starnix_uapi::vfs::FdEvents;
22use starnix_uapi::{errno, error, from_status_like_fdio, uapi, ucred};
23use std::sync::Arc;
24use zerocopy::IntoBytes;
25static READABLE_SIGNAL: zx::Signals =
26    zx::Signals::from_bits_retain(fio::FileSignal::READABLE.bits());
27static WRITABLE_SIGNAL: zx::Signals =
28    zx::Signals::from_bits_retain(fio::FileSignal::WRITABLE.bits());
29
30pub struct RemoteUnixDomainSocket {
31    client: fbinder::UnixDomainSocketSynchronousProxy,
32    event: zx::EventPair,
33    remote_creds: Arc<Credentials>,
34}
35
36impl RemoteUnixDomainSocket {
37    pub fn new(channel: zx::Channel, remote_creds: Arc<Credentials>) -> Result<Self, Errno> {
38        let client = fbinder::UnixDomainSocketSynchronousProxy::from_channel(channel);
39        let response = client
40            .get_event(
41                &fbinder::UnixDomainSocketGetEventRequest::default(),
42                zx::MonotonicInstant::INFINITE,
43            )
44            .map_err(|_| errno!(ECONNREFUSED))?
45            .map_err(|e: i32| from_status_like_fdio!(zx::Status::from_raw(e)))?;
46        let event = response.event.ok_or_else(|| errno!(ECONNREFUSED))?;
47        Ok(Self { client, event, remote_creds })
48    }
49
50    fn get_signals_from_events(events: FdEvents) -> zx::Signals {
51        let mut signals = zx::Signals::NONE;
52        if events.contains(FdEvents::POLLIN) {
53            signals |= READABLE_SIGNAL;
54        }
55        if events.contains(FdEvents::POLLOUT) {
56            signals |= WRITABLE_SIGNAL;
57        }
58        signals
59    }
60
61    fn get_events_from_signals(signals: zx::Signals) -> FdEvents {
62        let mut events = FdEvents::empty();
63        if signals.contains(READABLE_SIGNAL) {
64            events |= FdEvents::POLLIN;
65        }
66        if signals.contains(WRITABLE_SIGNAL) {
67            events |= FdEvents::POLLOUT;
68        }
69        events
70    }
71
72    /// Perform an action using the credentials of the remote task.
73    fn with_remote_creds<F, R>(&self, current_task: &CurrentTask, f: F) -> Result<R, Errno>
74    where
75        F: FnOnce() -> Result<R, Errno>,
76    {
77        current_task.override_creds(self.remote_creds.clone(), f)
78    }
79}
80
81impl SocketOps for RemoteUnixDomainSocket {
82    fn get_socket_info(&self) -> Result<(SocketDomain, SocketType, SocketProtocol), Errno> {
83        Ok((SocketDomain::Unix, SocketType::Datagram, SocketProtocol::from_raw(0)))
84    }
85
86    fn connect(
87        &self,
88        _socket: &SocketHandle,
89        _current_task: &CurrentTask,
90        _peer: SocketPeer,
91    ) -> Result<(), Errno> {
92        error!(EISCONN)
93    }
94
95    fn listen(&self, _socket: &Socket, _backlog: i32, _credentials: ucred) -> Result<(), Errno> {
96        error!(EOPNOTSUPP)
97    }
98
99    fn accept(&self, _socket: &Socket, _current_task: &CurrentTask) -> Result<SocketHandle, Errno> {
100        error!(EOPNOTSUPP)
101    }
102
103    fn bind(
104        &self,
105        _socket: &Socket,
106        _current_task: &CurrentTask,
107        _socket_address: SocketAddress,
108    ) -> Result<(), Errno> {
109        error!(EOPNOTSUPP)
110    }
111
112    fn read(
113        &self,
114        _socket: &Socket,
115        current_task: &CurrentTask,
116        data: &mut dyn OutputBuffer,
117        flags: SocketMessageFlags,
118    ) -> Result<MessageReadInfo, Errno> {
119        if self.client.is_closed().map_err(|_| errno!(ECONNREFUSED))? {
120            return error!(ECONNREFUSED);
121        }
122        let mut read_flags = fbinder::ReadFlags::empty();
123        if flags.contains(SocketMessageFlags::PEEK) {
124            read_flags |= fbinder::ReadFlags::PEEK;
125        }
126
127        let response = self
128            .client
129            .read(
130                &fbinder::UnixDomainSocketReadRequest {
131                    count: Some(data.available() as u64),
132                    flags: Some(read_flags),
133                    ..Default::default()
134                },
135                zx::MonotonicInstant::INFINITE,
136            )
137            .map_err(|_| errno!(ECONNREFUSED))?
138            .map_err(|e: i32| {
139                let status = zx::Status::from_raw(e);
140                if status == zx::Status::PEER_CLOSED {
141                    errno!(ECONNRESET)
142                } else {
143                    from_status_like_fdio!(status)
144                }
145            })?;
146
147        let written =
148            if let Some(received_data) = response.data { data.write(&received_data)? } else { 0 };
149
150        let mut file_handles: Vec<FileHandle> = vec![];
151        if let Some(handles) = response.handles {
152            // Use the remote task's credentials to create the remote_file object. This ensures
153            // that the SID associated to the fd is set to the correct value.
154            self.with_remote_creds(current_task, || {
155                for handle in handles {
156                    file_handles.push(new_remote_file(current_task, handle, OpenFlags::RDWR)?);
157                }
158                Ok(())
159            })?;
160        }
161        let ancillary_data = vec![AncillaryData::Unix(UnixControlData::Rights(file_handles))];
162
163        let message_length = response.data_original_length.unwrap_or(written as u64) as usize;
164
165        Ok(MessageReadInfo { bytes_read: written, message_length, address: None, ancillary_data })
166    }
167
168    fn write(
169        &self,
170        _socket: &Socket,
171        current_task: &CurrentTask,
172        data: &mut dyn InputBuffer,
173        _dest_address: &mut Option<SocketAddress>,
174        ancillary_data: &mut Vec<AncillaryData>,
175    ) -> Result<usize, Errno> {
176        if self.client.is_closed().map_err(|_| errno!(ECONNREFUSED))? {
177            return error!(ECONNREFUSED);
178        }
179
180        let mut handles: Vec<zx::NullableHandle> = vec![];
181        for data in ancillary_data {
182            match data {
183                AncillaryData::Unix(UnixControlData::Rights(file_handles)) => {
184                    // Access the served files with the credentials of the remote end.
185                    self.with_remote_creds(current_task, || {
186                        for file_handle in file_handles {
187                            let Some(handle) = file_handle.to_handle(current_task)? else {
188                                return error!(EINVAL);
189                            };
190                            handles.push(handle);
191                        }
192                        Ok(())
193                    })?;
194                }
195                _ => return error!(EINVAL),
196            }
197        }
198
199        let bytes = data.read_all()?;
200
201        let response = self
202            .client
203            .write(
204                fbinder::UnixDomainSocketWriteRequest {
205                    data: Some(bytes),
206                    handles: Some(handles),
207                    ..Default::default()
208                },
209                zx::MonotonicInstant::INFINITE,
210            )
211            .map_err(|_| errno!(ECONNREFUSED))?
212            .map_err(|e: i32| from_status_like_fdio!(zx::Status::from_raw(e)))?;
213
214        let written = response.actual_count.unwrap_or(0);
215        Ok(written as usize)
216    }
217
218    fn wait_async(
219        &self,
220        _socket: &Socket,
221        _current_task: &CurrentTask,
222        waiter: &Waiter,
223        events: FdEvents,
224        handler: EventHandler,
225    ) -> WaitCanceler {
226        let signal_handler = SignalHandler {
227            inner: SignalHandlerInner::ZxHandle(Self::get_events_from_signals),
228            event_handler: handler,
229            err_code: None,
230        };
231        let canceler = waiter
232            .wake_on_zircon_signals(
233                &self.event,
234                Self::get_signals_from_events(events),
235                signal_handler,
236            )
237            .unwrap();
238        WaitCanceler::new_port(canceler)
239    }
240
241    fn query_events(
242        &self,
243        _socket: &Socket,
244        _current_task: &CurrentTask,
245    ) -> Result<FdEvents, Errno> {
246        let signals = self
247            .event
248            .as_handle_ref()
249            .wait_one(zx::Signals::NONE, zx::MonotonicInstant::INFINITE_PAST)
250            .map_err(|e| from_status_like_fdio!(e))?;
251        Ok(Self::get_events_from_signals(signals))
252    }
253
254    fn shutdown(&self, _socket: &Socket, _how: SocketShutdownFlags) -> Result<(), Errno> {
255        Ok(())
256    }
257
258    fn close(&self, _current_task: &CurrentTask, _socket: &Socket) {
259        let _ = self.client.close(zx::MonotonicInstant::INFINITE);
260    }
261
262    fn getsockname(&self, _socket: &Socket) -> Result<SocketAddress, Errno> {
263        Ok(SocketAddress::default_for_domain(SocketDomain::Unix))
264    }
265
266    fn getpeername(&self, socket: &Socket) -> Result<SocketAddress, Errno> {
267        self.getsockname(socket)
268    }
269
270    fn setsockopt(
271        &self,
272        _socket: &Socket,
273        _current_task: &CurrentTask,
274        _level: u32,
275        _optname: u32,
276        _optval: SockOptValue,
277    ) -> Result<(), Errno> {
278        error!(EOPNOTSUPP)
279    }
280
281    fn getsockopt(
282        &self,
283        _socket: &Socket,
284        _current_task: &CurrentTask,
285        level: u32,
286        optname: u32,
287        _optlen: u32,
288    ) -> Result<Vec<u8>, Errno> {
289        if level != SOL_SOCKET {
290            return error!(EINVAL);
291        }
292        let data = match optname {
293            SO_LINGER => uapi::linger::default().as_bytes().to_vec(),
294            _ => return error!(EINVAL),
295        };
296
297        Ok(data)
298    }
299
300    fn to_handle(
301        &self,
302        _socket: &Socket,
303        _current_task: &CurrentTask,
304    ) -> Result<Option<zx::NullableHandle>, Errno> {
305        let (proxy, server) = zx::Channel::create();
306        self.client.clone(server.into()).map_err(|_| errno!(ECONNREFUSED))?;
307        Ok(Some(zx::NullableHandle::from(proxy).into()))
308    }
309}
310
311#[cfg(test)]
312mod tests {
313    use super::*;
314    use crate::testing::spawn_kernel_and_run;
315    use crate::vfs::socket::SocketFile;
316    use crate::vfs::{VecInputBuffer, VecOutputBuffer};
317    use fidl::endpoints::{DiscoverableProtocolMarker as _, RequestStream};
318    use fidl_fuchsia_unknown as funknown;
319    use fuchsia_async as fasync;
320    use futures::StreamExt;
321    use starnix_sync::{LockDepMutex, RemoteUnixDomainSocketStateLock};
322    use std::sync::Arc;
323
324    #[derive(Debug)]
325    struct Data {
326        bytes: Vec<u8>,
327        handles: Vec<zx::NullableHandle>,
328    }
329
330    impl Data {
331        fn try_clone(&mut self) -> Result<Self, zx::Status> {
332            let mut new_handles = vec![];
333            for handle in std::mem::take(&mut self.handles) {
334                let (new_handle, old_handle) = {
335                    let handle_type = handle.basic_info()?.object_type;
336                    match handle_type {
337                        zx::ObjectType::CHANNEL => {
338                            let channel = zx::Channel::from(handle);
339                            let client = funknown::CloneableSynchronousProxy::new(channel);
340                            let (proxy, server) = zx::Channel::create();
341                            let new_handle = client
342                                .clone(server.into())
343                                .map(|_| proxy.into())
344                                .map_err(|_| zx::Status::NOT_SUPPORTED);
345                            (new_handle, client.into_channel().into_handle())
346                        }
347                        _ => {
348                            let new_handle = handle.duplicate_handle(zx::Rights::SAME_RIGHTS);
349                            (new_handle, handle)
350                        }
351                    }
352                };
353                self.handles.push(old_handle);
354                new_handles.push(new_handle);
355            }
356            let new_handles =
357                new_handles.into_iter().collect::<Result<Vec<zx::NullableHandle>, zx::Status>>()?;
358            Ok(Self { bytes: self.bytes.clone(), handles: new_handles })
359        }
360    }
361
362    #[derive(Debug)]
363    struct UnixDomainSocketImplState {
364        _local_event: zx::EventPair,
365        remote_event: zx::EventPair,
366        buffer: Vec<Data>,
367    }
368
369    impl Default for UnixDomainSocketImplState {
370        fn default() -> Self {
371            let (_local_event, remote_event) = zx::EventPair::create();
372            Self { _local_event, remote_event, buffer: vec![] }
373        }
374    }
375
376    #[derive(Debug, Default)]
377    struct UnixDomainSocketImpl {
378        state: LockDepMutex<UnixDomainSocketImplState, RemoteUnixDomainSocketStateLock>,
379        close_on_read: bool,
380    }
381
382    impl UnixDomainSocketImpl {
383        fn read(
384            &self,
385            payload: fbinder::UnixDomainSocketReadRequest,
386        ) -> Result<fbinder::UnixDomainSocketReadResponse, zx::Status> {
387            if self.close_on_read {
388                return Err(zx::Status::PEER_CLOSED);
389            }
390            let Some(count) = payload.count else {
391                return Err(zx::Status::INVALID_ARGS);
392            };
393            let Some(flags) = payload.flags else {
394                return Err(zx::Status::INVALID_ARGS);
395            };
396            let mut state = self.state.lock();
397            if state.buffer.is_empty() {
398                return Err(zx::Status::SHOULD_WAIT);
399            }
400            let mut data = if flags.contains(fbinder::ReadFlags::PEEK) {
401                state.buffer[0].try_clone()?
402            } else {
403                state.buffer.remove(0)
404            };
405
406            if state.buffer.is_empty() {
407                state.remote_event.as_handle_ref().signal(READABLE_SIGNAL, WRITABLE_SIGNAL)?;
408            }
409
410            let actual_count = data.bytes.len() as u64;
411            data.bytes.truncate(count as usize);
412
413            Ok(fbinder::UnixDomainSocketReadResponse {
414                data: Some(data.bytes),
415                data_original_length: Some(actual_count),
416                handles: Some(data.handles),
417                ..Default::default()
418            })
419        }
420
421        fn write(
422            &self,
423            payload: fbinder::UnixDomainSocketWriteRequest,
424        ) -> Result<fbinder::UnixDomainSocketWriteResponse, zx::Status> {
425            let Some(bytes) = payload.data else {
426                return Err(zx::Status::INVALID_ARGS);
427            };
428            let actual_count = bytes.len() as u64;
429            let Some(handles) = payload.handles else {
430                return Err(zx::Status::INVALID_ARGS);
431            };
432            let mut state = self.state.lock();
433            state.buffer.push(Data { bytes, handles });
434            state
435                .remote_event
436                .as_handle_ref()
437                .signal(zx::Signals::NONE, READABLE_SIGNAL | WRITABLE_SIGNAL)?;
438            Ok(fbinder::UnixDomainSocketWriteResponse {
439                actual_count: Some(actual_count),
440                ..Default::default()
441            })
442        }
443
444        async fn serve(self: &Arc<Self>, channel: zx::Channel) {
445            let stream = fbinder::UnixDomainSocketRequestStream::from_channel(
446                fasync::Channel::from_channel(channel),
447            );
448            stream
449                .for_each_concurrent(None, |message| async {
450                    match message {
451                        Ok(fbinder::UnixDomainSocketRequest::GetEvent { responder, .. }) => {
452                            let state = self.state.lock();
453                            let event = state
454                                .remote_event
455                                .duplicate_handle(zx::Rights::SAME_RIGHTS)
456                                .expect("duplicate event");
457                            responder
458                                .send(Ok(fbinder::UnixDomainSocketGetEventResponse {
459                                    event: Some(event),
460                                    ..Default::default()
461                                }))
462                                .expect("respond");
463                        }
464                        Ok(fbinder::UnixDomainSocketRequest::Read {
465                            payload, responder, ..
466                        }) => {
467                            assert!(
468                                responder
469                                    .send(self.read(payload).map_err(|e| e.into_raw()))
470                                    .is_ok()
471                            );
472                        }
473                        Ok(fbinder::UnixDomainSocketRequest::Write {
474                            payload, responder, ..
475                        }) => {
476                            assert!(
477                                responder
478                                    .send(self.write(payload).as_ref().map_err(|e| e.into_raw()))
479                                    .is_ok()
480                            );
481                        }
482                        Ok(fbinder::UnixDomainSocketRequest::Query { responder }) => {
483                            assert!(
484                                responder
485                                    .send(fbinder::UnixDomainSocketMarker::PROTOCOL_NAME.as_bytes())
486                                    .is_ok()
487                            );
488                        }
489                        Ok(fbinder::UnixDomainSocketRequest::Clone { request, .. }) => {
490                            self.serve(request.into()).await;
491                        }
492                        Ok(fbinder::UnixDomainSocketRequest::Close { responder }) => {
493                            assert!(responder.send(Ok(())).is_ok());
494                        }
495                        _ => {
496                            return;
497                        }
498                    }
499                })
500                .await;
501        }
502    }
503
504    #[::fuchsia::test]
505    async fn test_remote_uds() {
506        let (client, server) = zx::Channel::create();
507        let handle = std::thread::spawn(|| {
508            let mut executor = fasync::LocalExecutor::default();
509            executor.run_singlethreaded(async move {
510                let uds_impl = UnixDomainSocketImpl::default();
511                Arc::new(uds_impl).serve(server).await;
512            });
513        });
514        spawn_kernel_and_run(async move |current_task| {
515            let original_file = new_remote_file(current_task, client.into(), OpenFlags::RDWR)
516                .expect("new_remote_file");
517            assert!(original_file.node().is_sock());
518            let file = new_remote_file(
519                current_task,
520                original_file.to_handle(current_task).expect("to_handle").expect("has_handle"),
521                OpenFlags::RDWR,
522            )
523            .expect("new_remote_file");
524            let ancillary_data =
525                vec![AncillaryData::Unix(UnixControlData::Rights(vec![original_file]))];
526            let socket_ops = file.downcast_file::<SocketFile>().unwrap();
527            let data = "HelloWorld";
528            let mut input_buffer = VecInputBuffer::new(data.as_bytes());
529            assert_eq!(
530                socket_ops.sendmsg(
531                    current_task,
532                    &file,
533                    &mut input_buffer,
534                    None,
535                    ancillary_data,
536                    SocketMessageFlags::empty()
537                ),
538                Ok(data.len())
539            );
540
541            let flags = SocketMessageFlags::CTRUNC
542                | SocketMessageFlags::TRUNC
543                | SocketMessageFlags::NOSIGNAL
544                | SocketMessageFlags::CMSG_CLOEXEC;
545
546            let mut buffer = VecOutputBuffer::new(1024);
547            let info = socket_ops
548                .recvmsg(&current_task, &file, &mut buffer, flags | SocketMessageFlags::PEEK, None)
549                .expect("recvmsg");
550
551            assert_eq!(info.ancillary_data.len(), 1);
552            assert_eq!(info.message_length, data.len());
553
554            let mut buffer = VecOutputBuffer::new(1024);
555            let info = socket_ops
556                .recvmsg(&current_task, &file, &mut buffer, flags, None)
557                .expect("recvmsg");
558
559            assert_eq!(info.ancillary_data.len(), 1);
560            assert_eq!(info.message_length, data.len());
561
562            let mut buffer = VecOutputBuffer::new(1024);
563            let err = socket_ops
564                .recvmsg(
565                    &current_task,
566                    &file,
567                    &mut buffer,
568                    flags | SocketMessageFlags::DONTWAIT,
569                    None,
570                )
571                .unwrap_err();
572            assert_eq!(err, errno!(EAGAIN));
573        })
574        .await;
575        handle.join().expect("join");
576    }
577
578    #[::fuchsia::test]
579    async fn test_remote_uds_peer_closed() {
580        let (client, server) = zx::Channel::create();
581        let handle = std::thread::spawn(move || {
582            let mut executor = fasync::LocalExecutor::default();
583            executor.run_singlethreaded(async move {
584                let uds_impl = UnixDomainSocketImpl { close_on_read: true, ..Default::default() };
585                Arc::new(uds_impl).serve(server).await;
586            });
587        });
588        spawn_kernel_and_run(async move |current_task| {
589            let file = new_remote_file(current_task, client.into(), OpenFlags::RDWR)
590                .expect("new_remote_file");
591            let socket_ops = file.downcast_file::<SocketFile>().unwrap();
592
593            let mut buffer = VecOutputBuffer::new(1024);
594            let err = socket_ops
595                .recvmsg(&current_task, &file, &mut buffer, SocketMessageFlags::empty(), None)
596                .unwrap_err();
597            assert_eq!(err, errno!(ECONNRESET));
598        })
599        .await;
600        handle.join().expect("join");
601    }
602}