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;
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(())
}
};
}
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;
assert_eq!(old_field_encoded_len, field_encoded_len);
return Err(TooLargeWorkOutput);
}
Ok(())
}
};
}
enum KernelStateWrapper {
Present(KernelState<OpenedFile>),
Absent {
num_fds: usize,
},
}
pub struct BoundedWorkOutput {
inner: CoreVmOutput,
kernel_state: KernelStateWrapper,
heap_page_mapper: PageMapper,
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)?;
max_encoded_len -= inner.vm_state.restart_host_call.encoded_size() as u32;
max_encoded_len += Some(u64::MAX).encoded_size() as u32;
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 })
}
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(),
},
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(())
}
pub fn take_restart_host_call(&mut self) -> Option<u64> {
self.inner.vm_state.restart_host_call.take()
}
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(())
}
pub fn take_kernel_state(&mut self) -> Result<KernelState<OpenedFile>, TooLargeWorkOutput> {
match self.kernel_state {
KernelStateWrapper::Present(ref mut kernel_state) => {
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)
},
KernelStateWrapper::Absent { .. } => unreachable!(),
}
}
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;
assert!(self.max_encoded_len <= max_report_elective_data());
self.kernel_state = KernelStateWrapper::Present(state);
},
KernelStateWrapper::Present(..) => unreachable!(),
}
}
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))
}
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> {
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!();
}
}