Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
199 changes: 156 additions & 43 deletions alioth/src/virtio/dev/vsock/uds_vsock.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
use std::collections::HashMap;
use std::fmt::Debug;
use std::fs;
use std::io::{BufRead, BufReader, BufWriter, ErrorKind, IoSlice, IoSliceMut, Read, Write};
use std::io::{self, BufRead, BufReader, BufWriter, ErrorKind, IoSlice, IoSliceMut, Read, Write};
use std::mem::size_of_val;
use std::num::Wrapping;
use std::os::fd::AsRawFd;
Expand Down Expand Up @@ -71,11 +71,29 @@ pub struct UdsVsock {
listener: UnixListener,
connections: HashMap<(u32, u32), Connection>,
ports: HashMap<Token, (u32, u32)>,
sockets: HashMap<Token, UnixStream>,
sockets: HashMap<Token, PendingConn>,
host_ports: HashMap<u32, u32>,
next_port: u32,
}

/// An accepted socket whose `CONNECT` request line has not been fully
/// received yet.
#[derive(Debug)]
struct PendingConn {
reader: BufReader<UnixStream>,
/// Bytes of the request line received so far.
msg: String,
}

/// Returns true if `e` means the host side of a connection is gone, which is
/// a normal event that must only affect that single connection.
fn is_conn_lost(e: &io::Error) -> bool {
matches!(
e.kind(),
ErrorKind::BrokenPipe | ErrorKind::ConnectionReset | ErrorKind::ConnectionAborted
)
}

fn get_buf_size(stream: &UnixStream) -> Result<usize> {
let mut buf_size = 0i32;
let mut arg_size = size_of_val(&buf_size) as libc::socklen_t;
Expand Down Expand Up @@ -106,42 +124,86 @@ impl UdsVsock {
}

fn create_socket(&mut self, registry: &Registry) -> Result<()> {
let (stream, _) = self.listener.accept()?;
stream.set_nonblocking(true)?;
let token = Token(stream.as_raw_fd() as usize);
registry.register(
&mut SourceFd(&stream.as_raw_fd()),
token,
Interest::READABLE,
)?;
self.sockets.insert(token, stream);
// The listener is registered edge-triggered, so drain the backlog.
// Otherwise a connection that arrives while another one is pending
// stalls until yet another client shows up.
loop {
let stream = match self.listener.accept() {
Ok((stream, _)) => stream,
Err(e) if e.kind() == ErrorKind::WouldBlock => break,
Err(e) if e.kind() == ErrorKind::Interrupted => continue,
Err(e) if is_conn_lost(&e) => {
log::debug!("{}: aborted connection: {e:?}", self.name);
continue;
}
Err(e) => return Err(e.into()),
};
stream.set_nonblocking(true)?;
let token = Token(stream.as_raw_fd() as usize);
registry.register(
&mut SourceFd(&stream.as_raw_fd()),
token,
Interest::READABLE,
)?;
let pending = PendingConn {
reader: BufReader::new(stream),
msg: String::new(),
};
self.sockets.insert(token, pending);
}
Ok(())
}

fn drop_socket(&self, socket: &UnixStream, registry: &Registry) -> Result<()> {
registry.deregister(&mut SourceFd(&socket.as_raw_fd()))?;
Ok(())
}

fn handle_conn_request<'m, Q, S>(
&mut self,
token: Token,
socket: UnixStream,
mut pending: PendingConn,
registry: &Registry,
rx_q: &mut Queue<'_, 'm, Q>,
irq_sender: &S,
) -> Result<()>
where
Q: VirtQueue<'m>,
S: IrqSender,
{
let mut msg = String::new();
let writer = socket.try_clone()?;
let mut reader = BufReader::new(socket);
// The socket is non-blocking, so a request line can arrive in pieces.
// Keep what has been received so far and wait for the next event
// instead of tearing down the connection.
match pending.reader.read_line(&mut pending.msg) {
Ok(_) => {}
Err(e) if e.kind() == ErrorKind::WouldBlock => {
self.sockets.insert(token, pending);
return Ok(());
}
Err(e) => return Err(e.into()),
}
if !pending.msg.ends_with('\n') {
if pending.msg.is_empty() {
log::debug!("{}: socket closed before any request", self.name);
} else {
log::warn!(
"{}: socket closed mid-request: {:?}",
self.name,
pending.msg
);
}
return self.drop_socket(pending.reader.get_ref(), registry);
}
let writer = pending.reader.get_ref().try_clone()?;
let buf_size = get_buf_size(&writer)?;
reader.read_line(&mut msg)?;
let port_str = msg.trim_start_matches("CONNECT ").trim_end();
let port_str = pending.msg.trim_start_matches("CONNECT ").trim_end();
let Ok(port) = port_str.parse::<u32>() else {
log::error!("{}: failed to parse port {port_str}", self.name);
return Ok(());
return self.drop_socket(pending.reader.get_ref(), registry);
};
let Some(host_port) = self.allocate_port() else {
log::error!("{}: failed to allocate port", self.name);
return Ok(());
return self.drop_socket(pending.reader.get_ref(), registry);
};
let hdr = VsockHeader {
src_cid: VSOCK_CID_HOST,
Expand All @@ -157,7 +219,7 @@ impl UdsVsock {
self.respond(&hdr, irq_sender, rx_q)?;
let conn = Connection {
state: ConnState::Requested,
reader,
reader: pending.reader,
writer: BufWriter::new(writer),
buf_alloc: buf_size as u32,
eof: false,
Expand Down Expand Up @@ -254,8 +316,20 @@ impl UdsVsock {
);
return Ok(());
};
writeln!(conn.writer, "OK {host_port}")?;
conn.writer.flush()?;
let acked = writeln!(conn.writer, "OK {host_port}").and_then(|_| conn.writer.flush());
match acked {
Ok(()) => {}
// The host hung up before the guest accepted the connection.
// `process_rx_data()` below turns this into an RST for the guest.
Err(e) if is_conn_lost(&e) => {
log::debug!(
"{}: host:{host_port} -> vm:{guest_port}: host closed before accept",
self.name
);
conn.eof = true;
}
Err(e) => return Err(e.into()),
}
conn.state = ConnState::Established {
fwd_cnt: Wrapping(0),
};
Expand Down Expand Up @@ -455,7 +529,7 @@ impl UdsVsock {
VsockOp::REQUEST => self.handle_tx_request(hdr, registry, irq_sender, rx_q),
VsockOp::RESPONSE => self.handle_tx_response(hdr, registry, rx_q, irq_sender),
VsockOp::RST => self.handle_tx_rst(hdr, registry),
VsockOp::RW => self.transfer_tx_data(hdr, body, readable),
VsockOp::RW => self.transfer_tx_data(hdr, body, readable, registry, rx_q, irq_sender),
VsockOp::CREDIT_UPDATE => {
log::info!(
"{name}: CREDIT_UPDATE: fwd_cnt: {}, buf_alloc: {}",
Expand Down Expand Up @@ -554,7 +628,13 @@ impl UdsVsock {
return Ok(());
}
let ConnState::Established { fwd_cnt } = conn.state else {
log::error!("{}: unexpected state {:?}", self.name, conn.state);
// Data can arrive before the guest accepts the connection. It
// stays buffered in the socket until then.
log::debug!(
"{}: host:{host_port} -> vm:{guest_port}: not ready, state {:?}",
self.name,
conn.state
);
return Ok(());
};
let mut hdr = VsockHeader {
Expand Down Expand Up @@ -652,17 +732,24 @@ impl UdsVsock {
Ok(())
}

fn transfer_tx_data(
fn transfer_tx_data<'m, Q, S>(
&mut self,
hdr: &VsockHeader,
body: &[u8],
buffers: &[IoSlice],
) -> Result<()> {
registry: &Registry,
rx_q: &mut Queue<'_, 'm, Q>,
irq_sender: &S,
) -> Result<()>
where
Q: VirtQueue<'m>,
S: IrqSender,
{
fn copy_to_conn(
buf: &[u8],
conn: &mut BufWriter<UnixStream>,
remain: &mut usize,
) -> Result<()> {
) -> io::Result<()> {
if let Some(b) = buf.get(..*remain) {
conn.write_all(b)?;
*remain = 0;
Expand All @@ -673,6 +760,28 @@ impl UdsVsock {
Ok(())
}

/// Writes up to `len` bytes of `body` and `buffers` to `conn`,
/// returning the number of bytes that were not covered by the input.
fn write_to_conn(
conn: &mut BufWriter<UnixStream>,
body: &[u8],
buffers: &[IoSlice],
len: usize,
) -> io::Result<usize> {
let mut remain = len;
if !body.is_empty() {
copy_to_conn(body, conn, &mut remain)?;
}
for buf in buffers {
if remain == 0 {
break;
}
copy_to_conn(buf, conn, &mut remain)?;
}
conn.flush()?;
Ok(remain)
}

let host_port = hdr.dst_port;
let guest_port = hdr.src_port;
let Some(conn) = self.connections.get_mut(&(host_port, guest_port)) else {
Expand All @@ -686,27 +795,30 @@ impl UdsVsock {
log::warn!("{}: invalid connection state {:?}", self.name, conn.state);
return Ok(());
};
let mut remain = hdr.len as usize;
if !body.is_empty() {
copy_to_conn(body, &mut conn.writer, &mut remain)?;
}
for buf in buffers {
if remain == 0 {
break;
match write_to_conn(&mut conn.writer, body, buffers, hdr.len as usize) {
Ok(0) => {}
Ok(remain) => {
log::error!("{}: missing {remain} bytes", self.name);
return error::InvalidBuffer.fail();
}
copy_to_conn(buf, &mut conn.writer, &mut remain)?;
}
if remain != 0 {
log::error!("{}: missing {remain} bytes", self.name);
return error::InvalidBuffer.fail();
// The host hung up. Reset this connection only, the rest of the
// device keeps running.
Err(e) if is_conn_lost(&e) => {
log::debug!(
"{}: vm:{guest_port} -> host:{host_port}: host closed",
self.name
);
conn.eof = true;
return self.process_rx_data(host_port, guest_port, registry, rx_q, irq_sender);
}
Err(e) => return Err(e.into()),
}
*fwd_cnt += hdr.len;
log::trace!(
"{}: vm:{guest_port} -> host:{host_port}: transferred {} bytes",
self.name,
hdr.len
);
conn.writer.flush()?;
Ok(())
}
}
Expand Down Expand Up @@ -837,8 +949,8 @@ impl VirtioMio for UdsVsock {
};
if token.0 == self.listener.as_raw_fd() as usize {
self.create_socket(registry)
} else if let Some(socket) = self.sockets.remove(&token) {
self.handle_conn_request(token, socket, rx_q, irq_sender)
} else if let Some(pending) = self.sockets.remove(&token) {
self.handle_conn_request(token, pending, registry, rx_q, irq_sender)
} else if let Some(port_pair) = self.ports.get(&token) {
let (host_port, guest_port) = port_pair.to_owned();
self.process_rx_data(host_port, guest_port, registry, rx_q, irq_sender)
Expand Down Expand Up @@ -887,7 +999,8 @@ impl VirtioMio for UdsVsock {
log::error!("{}: failed to deregister socket: {err}", self.name);
}
}
for (_, socket) in self.sockets.drain() {
for (_, pending) in self.sockets.drain() {
let socket = pending.reader.into_inner();
if let Err(err) = registry.deregister(&mut SourceFd(&socket.as_raw_fd())) {
log::error!("{}: failed to deregister socket: {err}", self.name);
}
Expand Down
Loading