Skip to content
Draft
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
325 changes: 305 additions & 20 deletions src/devices/src/virtio/console/device.rs
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
use std::cmp;
use std::io::Write;
use std::iter::zip;
use std::mem::{size_of, size_of_val};
use std::mem::size_of;
#[cfg(unix)]
use std::os::unix::io::{AsRawFd, RawFd};
use std::sync::Arc;
Expand All @@ -23,11 +23,78 @@ use crate::virtio::console::port::Port;
use crate::virtio::console::port_queue_mapping::{
QueueDirection, num_queues, port_id_to_queue_idx,
};
use crate::virtio::{InterruptTransport, PortDescription, VmmExitObserver};
use crate::virtio::descriptor_utils::Reader;
use crate::virtio::queue::{VIRTQ_DESC_F_NEXT, VIRTQ_DESC_F_WRITE};
use crate::virtio::{DescriptorChain, InterruptTransport, PortDescription, VmmExitObserver};

pub(crate) const CONTROL_RXQ_INDEX: usize = 2;
pub(crate) const CONTROL_TXQ_INDEX: usize = 3;

fn read_control_command(
mem: &GuestMemoryMmap,
head: DescriptorChain<'_>,
) -> Option<VirtioConsoleControl> {
let mut descriptor = Some(head.clone());
let mut descriptor_indices = Vec::new();
let mut readable_len = 0usize;

while let Some(current) = descriptor {
if descriptor_indices.contains(&current.index) {
log::error!("Console control descriptor chain contains a loop");
return None;
}
descriptor_indices.push(current.index);

if current.flags & !(VIRTQ_DESC_F_NEXT | VIRTQ_DESC_F_WRITE) != 0 {
log::error!("Console control transmit chain uses unsupported descriptor flags");
return None;
}

if current.is_write_only() {
log::error!("Console control transmit chain contains a writable descriptor");
return None;
}

readable_len = match readable_len.checked_add(current.len as usize) {
Some(len) => len,
None => {
log::error!("Console control descriptor length overflow");
return None;
}
};

let has_next = current.flags & VIRTQ_DESC_F_NEXT != 0;
descriptor = current.next_descriptor();
if has_next && descriptor.is_none() {
log::error!("Malformed console control descriptor chain");
return None;
}
}

if readable_len != size_of::<VirtioConsoleControl>() {
log::error!(
"Invalid console control payload length: expected {}, got {readable_len}",
size_of::<VirtioConsoleControl>()
);
return None;
}

let mut reader = match Reader::new(mem, head) {
Ok(reader) => reader,
Err(e) => {
log::error!("Invalid console control descriptor memory: {e}");
return None;
}
};
match reader.read_obj() {
Ok(command) => Some(command),
Err(e) => {
log::error!("Failed to read console control payload: {e}");
None
}
}
}

pub(crate) const AVAIL_FEATURES: u64 = (1 << uapi::VIRTIO_CONSOLE_F_SIZE as u64)
| (1 << uapi::VIRTIO_CONSOLE_F_MULTIPORT as u64)
| (1 << uapi::VIRTIO_F_VERSION_1 as u64);
Expand Down Expand Up @@ -177,23 +244,14 @@ impl Console {
while let Some(head) = control_tx.queue.pop(mem) {
raise_irq = true;

let cmd: VirtioConsoleControl = match mem.read_obj(head.addr) {
Ok(cmd) => cmd,
Err(e) => {
log::error!(
"Failed to read VirtioConsoleControl struct: {e:?}, struct len = {len}, head.len = {head_len}",
len = size_of::<VirtioConsoleControl>(),
head_len = head.len,
);
continue;
}
};
if let Err(e) = control_tx
.queue
.add_used(mem, head.index, size_of_val(&cmd) as u32)
{
let head_index = head.index;
let cmd = read_control_command(mem, head);
if let Err(e) = control_tx.queue.add_used(mem, head_index, 0) {
error!("failed to add used elements to the queue: {e:?}");
}
let Some(cmd) = cmd else {
continue;
};

log::trace!("VirtioConsoleControl cmd: {cmd:?}");
match cmd.event {
Expand All @@ -207,12 +265,17 @@ impl Console {
}
}
control_event::VIRTIO_CONSOLE_PORT_READY => {
let Some(port) = self.ports.get(cmd.id as usize) else {
log::warn!("Guest reported unknown console port {} ready", cmd.id);
continue;
};

if cmd.value != 1 {
log::error!("Port initialization failed: {cmd:?}");
continue;
}

if let Some(term) = self.ports[cmd.id as usize].terminal() {
if let Some(term) = port.terminal() {
self.control.mark_console_port(mem, cmd.id);
self.control.port_open(cmd.id, true);
let (cols, rows) = term.get_win_size();
Expand All @@ -224,13 +287,20 @@ impl Console {
self.control.port_open(cmd.id, true)
}

let name = self.ports[cmd.id as usize].name();
let name = port.name();
log::trace!("Port ready {id}: {name}", id = cmd.id);
if !name.is_empty() {
self.control.port_name(cmd.id, name)
}
}
control_event::VIRTIO_CONSOLE_PORT_OPEN => {
let Some(port_id) =
((cmd.id as usize) < self.ports.len()).then_some(cmd.id as usize)
else {
log::warn!("Guest reported unknown console port {} open", cmd.id);
continue;
};

let opened = match cmd.value {
0 => false,
1 => true,
Expand All @@ -249,7 +319,7 @@ impl Console {
continue;
}

ports_to_start.push(cmd.id as usize);
ports_to_start.push(port_id);
}
_ => log::warn!("Unknown console control event {:x}", cmd.event),
}
Expand Down Expand Up @@ -370,3 +440,218 @@ impl VmmExitObserver for Console {
log::trace!("Console on_vmm_exit finished");
}
}

#[cfg(test)]
mod tests {
use std::sync::Arc;

use utils::eventfd::EventFd;
use vm_memory::{Bytes, GuestAddress, GuestMemoryMmap};

use super::*;
use crate::legacy::DummyIrqChip;
use crate::virtio::console::console_control::Payload;
use crate::virtio::queue::tests::VirtQueue;
use crate::virtio::queue::{VIRTQ_DESC_F_NEXT, VIRTQ_DESC_F_WRITE};

const QUEUE_SIZE: u16 = 8;
const CONTROL_ADDR: u64 = 0x4000;

struct ProcessResult {
console: Console,
raised_irq: bool,
used_index: u16,
used_id: u32,
used_len: u32,
}

fn process_control_chain(
command: VirtioConsoleControl,
descriptors: &[(u64, u32, u16, u16)],
) -> ProcessResult {
let mem = GuestMemoryMmap::from_ranges(&[(GuestAddress(0), 0x10000)]).unwrap();
let queues = [0, 0x400, 0x800, 0xc00]
.map(|start| VirtQueue::new(GuestAddress(start), &mem, QUEUE_SIZE));

let command_bytes = command.as_slice();
let mut command_offset = 0;
for (index, &(addr, len, flags, next)) in descriptors.iter().enumerate() {
queues[CONTROL_TXQ_INDEX].dtable[index].set(addr, len, flags, next);
if flags & VIRTQ_DESC_F_WRITE == 0 && command_offset < command_bytes.len() {
let write_len = usize::min(len as usize, command_bytes.len() - command_offset);
if mem
.write(
&command_bytes[command_offset..command_offset + write_len],
GuestAddress(addr),
)
.is_ok()
{
command_offset += write_len;
}
}
}
queues[CONTROL_TXQ_INDEX].avail.ring[0].set(0);
queues[CONTROL_TXQ_INDEX].avail.idx.set(1);

let device_queues = queues
.iter()
.map(|queue| {
DeviceQueue::new(
queue.create_queue(),
Arc::new(EventFd::new(0).expect("create queue event")),
)
})
.collect();
let interrupt = InterruptTransport::new(DummyIrqChip::new().into(), "test".into())
.expect("create interrupt transport");
let mut console = Console::new(vec![PortDescription {
name: "test".into(),
input: None,
output: None,
terminal: None,
}])
.unwrap();
console
.activate(mem.clone(), interrupt, device_queues)
.unwrap();

let raised_irq = console.process_control_tx();
let used_index = queues[CONTROL_TXQ_INDEX].used.idx.get();
assert!(!console.process_control_tx());
assert_eq!(queues[CONTROL_TXQ_INDEX].used.idx.get(), used_index);

let used = queues[CONTROL_TXQ_INDEX].used.ring[0].get();
ProcessResult {
console,
raised_irq,
used_index,
used_id: used.id,
used_len: used.len,
}
}

#[test]
fn unknown_port_ids_are_completed_without_side_effects() {
for event in [
control_event::VIRTIO_CONSOLE_PORT_READY,
control_event::VIRTIO_CONSOLE_PORT_OPEN,
] {
for id in [1, u32::MAX] {
let command = VirtioConsoleControl {
id,
event,
value: 1,
};
let result = process_control_chain(
command,
&[(CONTROL_ADDR, size_of::<VirtioConsoleControl>() as u32, 0, 0)],
);

assert!(result.raised_irq);
assert_eq!(result.used_index, 1);
assert_eq!(result.used_id, 0);
assert_eq!(result.used_len, 0);
assert!(result.console.queues.iter().all(Option::is_some));
assert!(result.console.control.queue_pop().is_none());
}
}
}

#[test]
fn valid_control_chains_are_processed() {
let command = VirtioConsoleControl {
id: 0,
event: control_event::VIRTIO_CONSOLE_DEVICE_READY,
value: 1,
};
let chains = [
("contiguous", vec![(CONTROL_ADDR, 8, 0, 0)]),
(
"split",
vec![
(CONTROL_ADDR, 3, VIRTQ_DESC_F_NEXT, 1),
(CONTROL_ADDR + 0x100, 5, 0, 0),
],
),
];

for (name, chain) in chains {
let result = process_control_chain(command, &chain);

assert!(result.raised_irq, "{name}");
assert_eq!(result.used_index, 1, "{name}");
assert_eq!(result.used_id, 0, "{name}");
assert_eq!(result.used_len, 0, "{name}");
let Some(Payload::ConsoleControl(response)) = result.console.control.queue_pop() else {
panic!("expected port-add response");
};
let response_id = response.id;
let response_event = response.event;
assert_eq!(response_id, 0);
assert_eq!(response_event, control_event::VIRTIO_CONSOLE_PORT_ADD);
assert!(result.console.control.queue_pop().is_none());
}
}

#[test]
fn invalid_control_chains_are_completed_without_side_effects() {
let command = VirtioConsoleControl {
id: 0,
event: control_event::VIRTIO_CONSOLE_PORT_READY,
value: 1,
};
let chains = [
("short", vec![(CONTROL_ADDR, 7, 0, 0)]),
("trailing payload", vec![(CONTROL_ADDR, 9, 0, 0)]),
("write only", vec![(CONTROL_ADDR, 8, VIRTQ_DESC_F_WRITE, 0)]),
(
"mixed direction",
vec![
(CONTROL_ADDR, 4, VIRTQ_DESC_F_NEXT, 1),
(CONTROL_ADDR + 0x100, 4, VIRTQ_DESC_F_WRITE, 0),
],
),
("invalid first range", vec![(0x10000, 8, 0, 0)]),
(
"invalid second range",
vec![(CONTROL_ADDR, 4, VIRTQ_DESC_F_NEXT, 1), (0x10000, 4, 0, 0)],
),
(
"loop",
vec![
(CONTROL_ADDR, 4, VIRTQ_DESC_F_NEXT, 1),
(CONTROL_ADDR + 0x100, 4, VIRTQ_DESC_F_NEXT, 0),
],
),
("oversized length", vec![(CONTROL_ADDR, u32::MAX, 0, 0)]),
("unsupported flags", vec![(CONTROL_ADDR, 8, 4, 0)]),
];

for (name, chain) in chains {
let result = process_control_chain(command, &chain);

assert!(result.raised_irq, "{name}");
assert_eq!(result.used_index, 1, "{name}");
assert_eq!(result.used_id, 0, "{name}");
assert_eq!(result.used_len, 0, "{name}");
assert!(result.console.queues.iter().all(Option::is_some), "{name}");
assert!(result.console.control.queue_pop().is_none(), "{name}");
}
}

#[test]
fn malformed_control_head_has_no_side_effects() {
let command = VirtioConsoleControl {
id: 0,
event: control_event::VIRTIO_CONSOLE_PORT_READY,
value: 1,
};
let result =
process_control_chain(command, &[(CONTROL_ADDR, 8, VIRTQ_DESC_F_NEXT, QUEUE_SIZE)]);

assert!(!result.raised_irq);
assert_eq!(result.used_index, 0);
assert!(result.console.queues.iter().all(Option::is_some));
assert!(result.console.control.queue_pop().is_none());
}
}
Loading