corevm-engine 0.1.28

CoreVM engine that drives program execution either on the builder or CoreVM service side
Documentation
use crate::{OpenedFile, PageMapper, TooLargeWorkOutput};
use alloc::vec::Vec;
use codec::{Compact, CompactLen, ConstEncodedLen, Encode, MaxEncodedLen};
use corevm_host::{
	AudioMode, CoreVmOutput, KernelFd, PageNum, Range, ServiceMessage, VideoMode, VmOutput,
};
use jam_types::{max_report_elective_data, Hash, ServiceId};
use polkakernel::KernelState;

/// Set field value while checking that we haven't exceeded `max_report_elective_data`.
macro_rules! set_metered_field {
    ($setter:ident, $type:ty, $field:ident$(.$inner_field:ident)*) => {
        pub fn $setter(&mut self, value: $type) -> Result<(), TooLargeWorkOutput> {
            let max_encoded_len = self.max_encoded_len -
                self.inner.$field$(.$inner_field)*.encoded_size() as u32 +
                value.encoded_size() as u32;
            if max_encoded_len > max_report_elective_data() {
                return Err(TooLargeWorkOutput);
            }
            self.inner.$field$(.$inner_field)* = value;
            self.max_encoded_len = max_encoded_len;
            Ok(())
        }
    };
}

/// Update field value while checking that we haven't exceeded `max_report_elective_data`.
macro_rules! update_metered_field {
    ($updater:ident, $type:ty, $field:ident$(.$inner_field:ident)*) => {
        pub fn $updater(
            &mut self,
            update: impl FnOnce(&mut $type),
            restore: impl FnOnce(&mut $type),
        ) -> Result<(), TooLargeWorkOutput> {
            let old_max_encoded_len = self.max_encoded_len;
            let old_field_encoded_len = self.inner.$field$(.$inner_field)*.encoded_size() as u32;
            self.max_encoded_len -= old_field_encoded_len;
            update(&mut self.inner.$field$(.$inner_field)*);
            self.max_encoded_len += self.inner.$field$(.$inner_field)*.encoded_size() as u32;
            if self.max_encoded_len > max_report_elective_data() {
                restore(&mut self.inner.$field$(.$inner_field)*);
                self.max_encoded_len = old_max_encoded_len;
                let field_encoded_len = self.inner.$field$(.$inner_field)*.encoded_size() as u32;
                // If the lengths don't match, then `restore` doesn't actually reverse
                // the effects of `update`. This is a programming error!
                assert_eq!(old_field_encoded_len, field_encoded_len);
                return Err(TooLargeWorkOutput);
            }
            Ok(())
        }
    };
}

/// A wrapper for the kernel state that helps ensure that the state changes don't increase work
/// output size beyond limit.
///
/// To handle Linux syscall the state is extracted from the work output via `take_kernel_state`;
/// after handling the syscall it is returned back via `return_kernel_state`. Currently the size of
/// the state depends only on the number of file descriptors and Linux can allocate at most two
/// descriptors per syscall. Hence we account for two additional file descriptors before the
/// syscall is handled to be on the safe side.
enum KernelStateWrapper {
	/// The kernel state was returned back to the work output.
	Present(KernelState<OpenedFile>),
	/// The kernel state was taken from the work output to handle the syscall.
	Absent {
		/// The total number of file descriptors in the kernel.
		num_fds: usize,
	},
}

/// A wrapper for [`CoreVmOutput`] encoded size of which is bounded by [`max_report_elective_data`].
///
/// This struct ensures that the size of the work output never goes beyond this limit. Most of the
/// methods that modify it are implemented opportunistically, i.e. they modify the structure first
/// and then check the resulting encoded size and rollback changes if the limit is reached. This is
/// done under assumption that before reaching the limit methods are called multiple times,
/// thus it should be more efficient than pre-computing the resulting encoded size before modifying
/// the structure.
pub struct BoundedWorkOutput {
	inner: CoreVmOutput,
	kernel_state: KernelStateWrapper,
	heap_page_mapper: PageMapper,
	/// Upper boundary of the [`CoreVmOutput`]'s encoded size.
	///
	/// This is an upper boundary rather than the exact size because we track max. encoded size for
	/// some of the fields (for simplicity).
	///
	/// Includes `auth_output_len`.
	max_encoded_len: u32,
}

impl BoundedWorkOutput {
	pub fn try_from(
		mut inner: CoreVmOutput,
		kernel_state: KernelState<OpenedFile>,
		auth_output_len: u32,
		heap_page_range: Range,
	) -> Result<Self, TooLargeWorkOutput> {
		let mut max_encoded_len: u32 =
			inner.encoded_size().try_into().map_err(|_| TooLargeWorkOutput)?;
		// Adjust encoded size to account for fields that can be changed to any value without
		// causing `TooLargeWorkOutput` error. These fields are changed after the engine already
		// encountered another error and needs to update work output; we can't have double error in
		// this case.
		//
		// vm_state.restart_host_call
		max_encoded_len -= inner.vm_state.restart_host_call.encoded_size() as u32;
		max_encoded_len += Some(u64::MAX).encoded_size() as u32;
		// vm_output
		max_encoded_len -= inner.vm_output.encoded_size() as u32;
		max_encoded_len += VmOutput::max_encoded_len() as u32;
		max_encoded_len = max_encoded_len.checked_add(auth_output_len).ok_or(TooLargeWorkOutput)?;
		if max_encoded_len > max_report_elective_data() {
			return Err(TooLargeWorkOutput);
		}
		let kernel_state = KernelStateWrapper::Present(kernel_state);
		let mapped_heap_pages = core::mem::take(&mut inner.vm_state.mapped_heap_pages);
		let heap_page_mapper = PageMapper::new(mapped_heap_pages, heap_page_range);
		Ok(Self { inner, kernel_state, heap_page_mapper, max_encoded_len })
	}

	/// Convert to the underlying CoreVM output.
	///
	/// You shouldn't modify the resulting output in way that increases its encoded size to ensure
	/// that it still fits in [`max_report_elective_data`] bytes.
	pub fn into_inner(mut self) -> CoreVmOutput {
		self.inner.vm_state.mapped_heap_pages = self.heap_page_mapper.pages;
		self.inner.vm_state.kernel = match self.kernel_state {
			KernelStateWrapper::Present(kernel_state) => corevm_host::KernelState {
				fds: kernel_state
					.fds
					.into_iter()
					.map(|(fd, file)| {
						let kernel_fd =
							KernelFd { block_ref: file.block_ref, position: file.node.position() };
						(fd, kernel_fd)
					})
					.collect(),
			},
			// It's a programming error to not call `Self::return_kernel_state` before
			// `Engine::suspend`.
			KernelStateWrapper::Absent { .. } => unreachable!(),
		};
		self.inner
	}

	fn check_encoded_len(&self) -> Result<(), TooLargeWorkOutput> {
		if self.max_encoded_len > max_report_elective_data() {
			return Err(TooLargeWorkOutput);
		}
		Ok(())
	}

	/// Removes restart host-call index.
	///
	/// This method should be called to determine which host-call the guest was executing when it
	/// was suspended last time.
	pub fn take_restart_host_call(&mut self) -> Option<u64> {
		self.inner.vm_state.restart_host_call.take()
	}

	/// Sets restart host-call index.
	///
	/// This method should be called on each host-call fault.
	pub fn set_restart_host_call(&mut self, index: u64) {
		self.inner.vm_state.restart_host_call = Some(index);
	}

	set_metered_field!(set_output_video_mode, Option<VideoMode>, vm_state.video);
	set_metered_field!(set_output_audio_mode, Option<AudioMode>, vm_state.audio);
	update_metered_field!(
		update_processed_service_messages,
		Vec<ServiceMessage>,
		processed_service_messages
	);
	update_metered_field!(update_outgoing_messages, Vec<(ServiceId, Vec<u8>)>, outgoing_messages);

	pub fn insert_updated_page(
		&mut self,
		page: PageNum,
		imported_page_hash: Option<Hash>,
	) -> Result<(), TooLargeWorkOutput> {
		let old_max_encoded_len = self.max_encoded_len;
		let page_inserted = {
			let old_encoded_len = self.inner.updated_pages.encoded_size() as u32;
			let inserted = self.inner.updated_pages.insert(page, [0; 32]).is_none();
			if inserted {
				self.max_encoded_len -= old_encoded_len;
				self.max_encoded_len += self.inner.updated_pages.encoded_size() as u32;
			}
			inserted
		};
		let resident_page_inserted = {
			let inserted = !self.inner.vm_state.resident_pages.contains_index(page.0);
			if inserted {
				let old_encoded_len = self.inner.vm_state.resident_pages.encoded_size() as u32;
				self.inner.vm_state.resident_pages.insert(Range::new(page.0, page.0 + 1));
				let new_encoded_len = self.inner.vm_state.resident_pages.encoded_size() as u32;
				self.max_encoded_len -= old_encoded_len;
				self.max_encoded_len += new_encoded_len;
			}
			inserted
		};
		let touched_imported_page_inserted = imported_page_hash
			.map(|hash| {
				let old_encoded_len = self.inner.touched_imported_pages.encoded_size() as u32;
				let inserted = self.inner.touched_imported_pages.insert(page, hash).is_none();
				if inserted {
					let new_encoded_len = self.inner.touched_imported_pages.encoded_size() as u32;
					self.max_encoded_len -= old_encoded_len;
					self.max_encoded_len += new_encoded_len;
				}
				inserted
			})
			.unwrap_or(false);
		self.check_encoded_len().inspect_err(|_| {
			if page_inserted {
				self.inner.updated_pages.remove(&page);
			}
			if resident_page_inserted {
				self.inner.vm_state.resident_pages.remove(&Range::new(page.0, page.0 + 1));
			}
			if touched_imported_page_inserted {
				self.inner.touched_imported_pages.remove(&page);
			}
			self.max_encoded_len = old_max_encoded_len;
		})?;
		Ok(())
	}

	/// Extracts [`KernelState`] from the work output to handle a syscall.
	///
	/// The call to this method *must* be paired with a call to [`Self::return_kernel_state`] after
	/// the syscall is handled.
	pub fn take_kernel_state(&mut self) -> Result<KernelState<OpenedFile>, TooLargeWorkOutput> {
		match self.kernel_state {
			KernelStateWrapper::Present(ref mut kernel_state) => {
				// Here we reserve two file descriptors in case Linux system call tries to create
				// them.
				let num_fds = kernel_state.fds.len();
				self.max_encoded_len -= vec_map_encoded_len::<u32, KernelFd>(num_fds) as u32;
				self.max_encoded_len +=
					vec_map_encoded_len::<u32, KernelFd>(num_fds.saturating_add(2)) as u32;
				if self.max_encoded_len > max_report_elective_data() {
					return Err(TooLargeWorkOutput);
				}
				let kernel_state = core::mem::take(kernel_state);
				self.kernel_state = KernelStateWrapper::Absent { num_fds };
				Ok(kernel_state)
			},
			// This is a programming error.
			KernelStateWrapper::Absent { .. } => unreachable!(),
		}
	}

	/// Inserts previously extracted [`KernelState`] back into the work output.
	pub fn return_kernel_state(&mut self, state: KernelState<OpenedFile>) {
		match self.kernel_state {
			KernelStateWrapper::Absent { num_fds } => {
				self.max_encoded_len -=
					vec_map_encoded_len::<u32, KernelFd>(num_fds.saturating_add(2)) as u32;
				self.max_encoded_len +=
					vec_map_encoded_len::<u32, KernelFd>(state.fds.len()) as u32;
				// If this assertion fails, then either the kernel created more than two file
				// descriptors in one system call (which shouldn't be possible) or work output
				// was modified without returning kernel state first (which is a programming
				// error).
				assert!(self.max_encoded_len <= max_report_elective_data());
				self.kernel_state = KernelStateWrapper::Present(state);
			},
			// This is a programming error.
			KernelStateWrapper::Present(..) => unreachable!(),
		}
	}

	/// Maps specified number of memory pages.
	///
	/// Returns the numbers of mapped pages or `None` if the mapping failed.
	pub fn map_pages(&mut self, num_pages: u32) -> Result<Option<Range>, TooLargeWorkOutput> {
		let old_max_encoded_len = self.max_encoded_len;
		let old_encoded_len = self.heap_page_mapper.pages.encoded_size() as u32;
		let Some(page_range) = self.heap_page_mapper.map(num_pages) else {
			return Ok(None);
		};
		let new_encoded_len = self.heap_page_mapper.pages.encoded_size() as u32;
		self.max_encoded_len -= old_encoded_len;
		self.max_encoded_len += new_encoded_len;
		self.check_encoded_len().inspect_err(|_| {
			self.heap_page_mapper.unmap(page_range.start, page_range.end);
			self.max_encoded_len = old_max_encoded_len;
		})?;
		Ok(Some(page_range))
	}

	/// Unmaps specified memory pages.
	pub fn unmap_pages(&mut self, start_page: u32, end_page: u32) {
		self.max_encoded_len -= self.heap_page_mapper.pages.encoded_size() as u32;
		self.max_encoded_len -= self.inner.vm_state.resident_pages.encoded_size() as u32;
		self.heap_page_mapper.unmap(start_page, end_page);
		self.inner.vm_state.resident_pages.remove(&Range::new(start_page, end_page));
		self.max_encoded_len += self.heap_page_mapper.pages.encoded_size() as u32;
		self.max_encoded_len += self.inner.vm_state.resident_pages.encoded_size() as u32;
	}

	pub fn heap_page_mapper(&self) -> &PageMapper {
		&self.heap_page_mapper
	}
}

impl core::ops::Deref for BoundedWorkOutput {
	type Target = CoreVmOutput;

	fn deref(&self) -> &Self::Target {
		&self.inner
	}
}

fn vec_map_encoded_len<K: ConstEncodedLen, V: ConstEncodedLen>(len: usize) -> usize {
	(K::max_encoded_len() + V::max_encoded_len()) * len + Compact::<u64>::compact_len(&(len as u64))
}

#[cfg(test)]
mod tests {
	use super::*;
	use crate::testing::default_corevm_output;
	use jam_types::max_exports;

	impl BoundedWorkOutput {
		pub fn check_len(&self) -> Result<(), usize> {
			// We don't account for auth_output_len here.
			let len = self.inner.encoded_size();
			if len > max_report_elective_data() as usize {
				return Err(len);
			}
			Ok(())
		}

		pub fn reset_updated_pages(&mut self) {
			self.max_encoded_len -= self.inner.updated_pages.encoded_size() as u32;
			self.max_encoded_len -= self.inner.vm_state.resident_pages.encoded_size() as u32;
			self.inner.updated_pages.clear();
			self.inner.vm_state.resident_pages.clear();
			self.max_encoded_len += self.inner.updated_pages.encoded_size() as u32;
			self.max_encoded_len += self.inner.vm_state.resident_pages.encoded_size() as u32;
		}
	}

	#[test]
	fn insert_updated_page_works() {
		let mut work_output = BoundedWorkOutput::try_from(
			default_corevm_output(),
			Default::default(),
			0,
			Range::new(1000, 20000),
		)
		.unwrap();
		for i in 0..max_exports() + 1 {
			if work_output.insert_updated_page(PageNum(1000 + i), None).is_err() {
				assert_ne!(i, 0);
				let actual_encoded_size = work_output.into_inner().encoded_size();
				assert!(
					actual_encoded_size <= max_report_elective_data() as usize,
					"actual_encoded_size = {actual_encoded_size}, max_report_elective_data = {}",
					max_report_elective_data()
				);
				return;
			}
		}
		unreachable!();
	}
}