1use 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
19pub struct VsockSocket {
23 inner: LockDepMutex<VsockSocketInner, VsockSocketInnerLock>,
24}
25
26struct VsockSocketInner {
27 address: Option<SocketAddress>,
29
30 waiters: WaitQueue,
32
33 state: VsockSocketState,
35}
36
37enum VsockSocketState {
38 Disconnected,
40
41 Listening(AcceptQueue),
43
44 Connected { file: FileHandle, peer_addr: SocketAddress },
46
47 Closed,
49}
50
51fn downcast_socket_to_vsock(socket: &Socket) -> &VsockSocket {
52 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 fn lock(&self) -> LockDepGuard<'_, VsockSocketInner> {
73 self.inner.lock()
74 }
75}
76
77impl SocketOps for VsockSocket {
78 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 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 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 true,
266 )?;
267 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 ¤t_task,
327 SocketDomain::Vsock,
328 SocketType::Stream,
329 SocketProtocol::default(),
330 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(¤t_task, VSOCK_PORT, &listen_socket)
336 .expect("Failed to bind socket.");
337 listen_socket.listen(¤t_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(¤t_task, fs2, OpenFlags::RDWR | OpenFlags::NONBLOCK)
343 .unwrap();
344 listen_socket
345 .downcast_socket::<VsockSocket>()
346 .unwrap()
347 .remote_connection(&listen_socket, ¤t_task, remote)
348 .unwrap();
349
350 let server_socket = listen_socket.accept(¤t_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(¤t_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(¤t_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(¤t_task);
373 listen_socket.close(¤t_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 ¤t_task,
385 SocketDomain::Vsock,
386 SocketType::Stream,
387 SocketProtocol::default(),
388 false,
389 )
390 .expect("Failed to create socket.");
391 let remote =
392 create_fuchsia_pipe(¤t_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(¤t_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 std::thread::sleep(std::time::Duration::from_secs(2));
419
420 socket_file.write(¤t_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(¤t_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 ¤t_task,
439 SocketDomain::Vsock,
440 SocketType::Stream,
441 SocketProtocol::default(),
442 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 ¤t_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(¤t_task),
462 Ok(FdEvents::POLLOUT | FdEvents::POLLWRNORM)
463 );
464
465 let epoll_object = EpollFileObject::new_file(¤t_task);
466 let epoll_file = epoll_object.downcast_file::<EpollFileObject>().unwrap();
467 let event = EpollEvent::new(FdEvents::POLLIN, 0);
468 epoll_file.add(¤t_task, &socket, &epoll_object, event).expect("poll_file.add");
469
470 let fds = epoll_file.wait(¤t_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(¤t_task),
477 Ok(FdEvents::POLLOUT
478 | FdEvents::POLLWRNORM
479 | FdEvents::POLLIN
480 | FdEvents::POLLRDNORM)
481 );
482 let fds = epoll_file.wait(¤t_task, 1, zx::MonotonicInstant::ZERO).expect("wait");
483 assert_eq!(fds.len(), 1);
484
485 assert_eq!(socket.read(¤t_task, &mut VecOutputBuffer::new(64)).expect("read"), 1);
486
487 assert_eq!(
488 socket.query_events(¤t_task),
489 Ok(FdEvents::POLLOUT | FdEvents::POLLWRNORM)
490 );
491 let fds = epoll_file.wait(¤t_task, 1, zx::MonotonicInstant::ZERO).expect("wait");
492 assert!(fds.is_empty());
493 })
494 .await;
495 }
496}