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
63 changes: 63 additions & 0 deletions crates/machine/src/cow_ram.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@ pub struct CowRam {
len: usize,
mask: u64,
epoch: u64,
device_written_pages: Vec<u32>,
device_written_seen: Vec<u64>,
}

impl CowRam {
Expand All @@ -35,6 +37,8 @@ impl CowRam {
len: ram_size as usize + 8,
mask: ram_size - 1,
epoch: next_epoch(),
device_written_pages: Vec::new(),
device_written_seen: Vec::new(),
}
}

Expand All @@ -53,6 +57,8 @@ impl CowRam {
len,
mask: ram_size - 1,
epoch: next_epoch(),
device_written_pages: Vec::new(),
device_written_seen: Vec::new(),
}
}

Expand All @@ -70,6 +76,8 @@ impl CowRam {
len: logical_len,
mask: ram_size - 1,
epoch: next_epoch(),
device_written_pages: Vec::new(),
device_written_seen: Vec::new(),
}
}

Expand All @@ -80,6 +88,8 @@ impl CowRam {
len: self.len,
mask: self.mask,
epoch: next_epoch(),
device_written_pages: Vec::new(),
device_written_seen: Vec::new(),
}
}

Expand All @@ -88,6 +98,34 @@ impl CowRam {
self.epoch
}

pub fn note_device_write(&mut self, index: usize, len: usize) {
if len == 0 {
return;
}
let first_page = index / PAGE_SIZE;
let last_page = (index + len - 1) / PAGE_SIZE;

for page in first_page..=last_page {
let word = page / 64;
let bit = 1u64 << (page % 64);
if word >= self.device_written_seen.len() {
self.device_written_seen.resize(word + 1, 0);
}
if self.device_written_seen[word] & bit == 0 {
self.device_written_seen[word] |= bit;
self.device_written_pages.push(page as u32);
}
}
}

pub fn drain_device_written_pages(&mut self, mut visit: impl FnMut(usize)) {
for page in self.device_written_pages.drain(..) {
let page = page as usize;
self.device_written_seen[page / 64] &= !(1u64 << (page % 64));
visit(page);
}
}

#[inline(always)]
pub fn page_ptr(&self, page: usize) -> *const u8 {
self.page_ref(page).as_ptr()
Expand Down Expand Up @@ -305,6 +343,31 @@ mod tests {

const WASM32_ISIZE_MAX: usize = (1usize << 31) - 1;

#[test]
fn device_writes_are_listed_once_per_page_until_drained() {
let mut ram = CowRam::new(1 << 20);
ram.note_device_write(10, 4);
ram.note_device_write(20, 4);
ram.note_device_write(PAGE_SIZE - 2, 4);

let mut drained = Vec::new();
ram.drain_device_written_pages(|page| drained.push(page));
assert_eq!(drained, vec![0, 1]);

let mut again = Vec::new();
ram.drain_device_written_pages(|page| again.push(page));
assert!(again.is_empty(), "a drain left pages behind: {again:?}");

ram.note_device_write(5, 1);
let mut rewritten = Vec::new();
ram.drain_device_written_pages(|page| rewritten.push(page));
assert_eq!(
rewritten,
vec![0],
"a drained page could not be recorded again"
);
}

#[test]
fn new_allocates_page_padded_and_never_grows() {
for mb in [1u64, 64, 256, 512, 1024] {
Expand Down
81 changes: 80 additions & 1 deletion crates/machine/src/machine_bus.rs
Original file line number Diff line number Diff line change
Expand Up @@ -370,8 +370,14 @@ impl SystemBus for MachineBus {
0
}

#[inline]
#[inline(always)]
fn read_halfword(&mut self, address: u64) -> u16 {
if address >= RAM_BASE && address + 1 < RAM_BASE + self.ram.len() as u64 {
let index = (address - RAM_BASE) as usize;

return self.ram.read_u16(index);
}

u16::from_le_bytes([self.read_byte(address), self.read_byte(address + 1)])
}

Expand Down Expand Up @@ -592,6 +598,13 @@ impl SystemBus for MachineBus {
Some(self.ram.page_mut_ptr(page))
}

fn drain_device_written_pages(&mut self, visit: &mut dyn FnMut(u64)) -> bool {
let first_ram_page = RAM_BASE >> 12;
self.ram
.drain_device_written_pages(|page| visit(first_ram_page + page as u64));
true
}

#[inline(always)]
fn ram_epoch(&self) -> u64 {
self.ram.epoch()
Expand Down Expand Up @@ -829,6 +842,72 @@ mod tests {
);
}

const JUMP_TO_SELF: u32 = 0x0000_006f;
const FENCE_I: u32 = 0x0000_100f;
const FENCE_I_PAGE: u64 = RAM_BASE + 0x1000;

fn load_x5_with(value: u32) -> u32 {
(value << 20) | (5 << 7) | 0x13
}

fn load_words(bus: &mut MachineBus, address: u64, words: &[u32]) {
let bytes: Vec<u8> = words.iter().flat_map(|word| word.to_le_bytes()).collect();
bus.load_ram(address - RAM_BASE, &bytes);
}

fn run_from(hart: &mut Hart, bus: &mut MachineBus, pc: u64) {
hart.regs.pc = pc;
hart.run(bus, 16);
}

fn machine_with_code_decoded_at_ram_base() -> (Hart, MachineBus) {
const RAM_SIZE: u64 = 1 << 20;

let mut bus = MachineBus::new(RAM_SIZE, CowRam::new(RAM_SIZE));
let mut hart = Hart::new(RAM_BASE);
load_words(&mut bus, RAM_BASE, &[load_x5_with(1), JUMP_TO_SELF]);
load_words(&mut bus, FENCE_I_PAGE, &[FENCE_I, JUMP_TO_SELF]);

run_from(&mut hart, &mut bus, RAM_BASE);
assert_eq!(hart.regs.read(5), 1);
assert!(
hart.blocks.lookup(RAM_BASE).is_some(),
"nothing was decoded"
);

(hart, bus)
}

#[test]
fn code_a_device_overwrote_is_decoded_again_after_fence_i() {
let (mut hart, mut bus) = machine_with_code_decoded_at_ram_base();

let mask = bus.ram_mask;
RamView::new(&mut bus.ram, mask).write_u32(RAM_BASE, load_x5_with(2));
run_from(&mut hart, &mut bus, FENCE_I_PAGE);
run_from(&mut hart, &mut bus, RAM_BASE);

assert_eq!(
hart.regs.read(5),
2,
"a block decoded before a device overwrote its page survived FENCE.I, \
so the guest ran code that is no longer in memory"
);
}

#[test]
fn fence_i_keeps_code_nothing_overwrote() {
let (mut hart, mut bus) = machine_with_code_decoded_at_ram_base();

run_from(&mut hart, &mut bus, FENCE_I_PAGE);

assert!(
hart.blocks.lookup(RAM_BASE).is_some(),
"FENCE.I dropped code nothing had overwritten; a JIT fences thousands \
of times per process and would re-decode its working set on each"
);
}

#[test]
fn claim_still_reports_a_line_that_is_asserted() {
let mut bus = bus_with_one_console_byte();
Expand Down
5 changes: 5 additions & 0 deletions crates/machine/src/virtio/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,7 @@ impl<'a> RamView<'a> {
}
let idx = self.idx(physical_address);
self.ram.write_u8(idx, val);
self.ram.note_device_write(idx, 1);
}

pub fn write_u16(&mut self, physical_address: u64, val: u16) {
Expand All @@ -149,6 +150,7 @@ impl<'a> RamView<'a> {
}
let idx = self.idx(physical_address);
self.ram.write_u16(idx, val);
self.ram.note_device_write(idx, 2);
}

pub fn write_u32(&mut self, physical_address: u64, val: u32) {
Expand All @@ -158,6 +160,7 @@ impl<'a> RamView<'a> {
}
let idx = self.idx(physical_address);
self.ram.write_u32(idx, val);
self.ram.note_device_write(idx, 4);
}

pub fn write_u64(&mut self, physical_address: u64, val: u64) {
Expand All @@ -167,6 +170,7 @@ impl<'a> RamView<'a> {
}
let idx = self.idx(physical_address);
self.ram.write_u64(idx, val);
self.ram.note_device_write(idx, 8);
}

pub fn read_bytes(&self, physical_address: u64, buf: &mut [u8]) {
Expand All @@ -188,6 +192,7 @@ impl<'a> RamView<'a> {
None => {
let idx = self.idx(physical_address);
self.ram.write_from(idx, buf);
self.ram.note_device_write(idx, buf.len());
}
}
}
Expand Down
Loading
Loading