Skip to content
Open
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
145 changes: 134 additions & 11 deletions crates/environ/src/fact/trampoline.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1761,6 +1761,8 @@ impl<'a, 'b> Compiler<'a, 'b> {
Transcode::Latin1ToUtf16
};
let transcode = self.transcoder(src, &dst, op);
self.validate_transcode_buffer(src.opts, &src.ptr, &src.len, src_enc);
self.validate_transcode_buffer(dst.opts, &dst.ptr, &src.len, dst_enc);
self.instruction(LocalGet(src.ptr.idx));
self.instruction(LocalGet(src.len.idx));
self.instruction(LocalGet(dst.ptr.idx));
Expand Down Expand Up @@ -1832,6 +1834,8 @@ impl<'a, 'b> Compiler<'a, 'b> {
FE::Utf8 => unreachable!(),
};
let transcode = self.transcoder(src, &dst, op);
self.validate_transcode_buffer(src.opts, &src.ptr, &src.len, src_enc);
self.validate_transcode_buffer(dst.opts, &dst.ptr, &dst_byte_len, FE::Utf8);
self.instruction(LocalGet(src.ptr.idx));
self.instruction(LocalGet(src.len.idx));
self.instruction(LocalGet(dst.ptr.idx));
Expand Down Expand Up @@ -1892,17 +1896,31 @@ impl<'a, 'b> Compiler<'a, 'b> {
self.ptr_shl(src_mem_opts);
}
self.ptr_add(src_mem_opts);
let rest_src_ptr = self.local_set_new_tmp(src_mem_opts.ptr());
self.instruction(LocalGet(src.len.idx));
self.instruction(LocalGet(src_len_tmp.idx));
self.ptr_sub(src_mem_opts);
let rest_src_len = self.local_set_new_tmp(src_mem_opts.ptr());
self.instruction(LocalGet(dst.ptr.idx));
self.instruction(LocalGet(dst.len.idx));
self.ptr_add(dst_mem_opts);
let rest_dst_ptr = self.local_set_new_tmp(dst_mem_opts.ptr());
self.instruction(LocalGet(dst_byte_len.idx));
self.instruction(LocalGet(dst.len.idx));
self.ptr_sub(dst_mem_opts);
let rest_dst_len = self.local_set_new_tmp(dst_mem_opts.ptr());
self.validate_transcode_buffer(src.opts, &rest_src_ptr, &rest_src_len, src_enc);
self.validate_transcode_buffer(dst.opts, &rest_dst_ptr, &rest_dst_len, FE::Utf8);
self.instruction(LocalGet(rest_src_ptr.idx));
self.instruction(LocalGet(rest_src_len.idx));
self.instruction(LocalGet(rest_dst_ptr.idx));
self.instruction(LocalGet(rest_dst_len.idx));
self.instruction(I32Const(0)); // first_pass = false
self.instruction(Call(transcode.as_u32()));
self.free_temp_local(rest_src_ptr);
self.free_temp_local(rest_src_len);
self.free_temp_local(rest_dst_ptr);
self.free_temp_local(rest_dst_len);

// Add the second result, the amount of destination units encoded,
// to `dst_len` so it's an accurate reflection of the final size of
Expand Down Expand Up @@ -2014,6 +2032,8 @@ impl<'a, 'b> Compiler<'a, 'b> {
};

let transcode = self.transcoder(src, &dst, Transcode::Utf8ToUtf16);
self.validate_transcode_buffer(src.opts, &src.ptr, &src.len, FE::Utf8);
self.validate_transcode_buffer(dst.opts, &dst.ptr, &src.len, FE::Utf16);
self.instruction(LocalGet(src.ptr.idx));
self.instruction(LocalGet(src.len.idx));
self.instruction(LocalGet(dst.ptr.idx));
Expand Down Expand Up @@ -2101,6 +2121,8 @@ impl<'a, 'b> Compiler<'a, 'b> {
let src_byte_len = self.local_set_new_tmp(src_mem_opts.ptr());

let transcode = self.transcoder(src, &dst, Transcode::Utf16ToCompactProbablyUtf16);
self.validate_transcode_buffer(src.opts, &src.ptr, &src.len, FE::Utf16);
self.validate_transcode_buffer(dst.opts, &dst.ptr, &src.len, FE::Utf16);
self.instruction(LocalGet(src.ptr.idx));
self.instruction(LocalGet(src.len.idx));
self.instruction(LocalGet(dst.ptr.idx));
Expand Down Expand Up @@ -2194,6 +2216,8 @@ impl<'a, 'b> Compiler<'a, 'b> {
};
let transcode_latin1 = self.transcoder(src, &dst, latin1);
let transcode_utf16 = self.transcoder(src, &dst, utf16);
self.validate_transcode_buffer(src.opts, &src.ptr, &src.len, src_enc);
self.validate_transcode_buffer(dst.opts, &dst.ptr, &src.len, FE::Latin1);
self.instruction(LocalGet(src.ptr.idx));
self.instruction(LocalGet(src.len.idx));
self.instruction(LocalGet(dst.ptr.idx));
Expand Down Expand Up @@ -2265,14 +2289,22 @@ impl<'a, 'b> Compiler<'a, 'b> {
self.ptr_shl(src_mem_opts);
}
self.ptr_add(src_mem_opts);
let rest_src_ptr = self.local_set_new_tmp(src_mem_opts.ptr());
self.instruction(LocalGet(src.len.idx));
self.instruction(LocalGet(src_len_tmp.idx));
self.ptr_sub(src_mem_opts);
let rest_src_len = self.local_set_new_tmp(src_mem_opts.ptr());
self.validate_transcode_buffer(src.opts, &rest_src_ptr, &rest_src_len, src_enc);
self.validate_transcode_buffer(dst.opts, &dst.ptr, &src.len, FE::Utf16);
self.instruction(LocalGet(rest_src_ptr.idx));
self.instruction(LocalGet(rest_src_len.idx));
self.instruction(LocalGet(dst.ptr.idx));
self.convert_src_len_to_dst(src.len.idx, src_mem_opts.ptr(), dst_mem_opts.ptr());
self.instruction(LocalGet(dst.len.idx));
self.instruction(Call(transcode_utf16.as_u32()));
self.instruction(LocalSet(dst.len.idx));
self.free_temp_local(rest_src_ptr);
self.free_temp_local(rest_src_len);

// If the returned number of code units written to the destination
// is not equal to the size of the allocation then the allocation is
Expand Down Expand Up @@ -2330,6 +2362,86 @@ impl<'a, 'b> Compiler<'a, 'b> {
self.instruction(End);
}

/// Traps unless the `len` code units of `enc` at `ptr` are in-bounds, by
/// loading the last code unit.
///
/// Buffers are already validated when they're created, but host
/// transcoders turn them into raw slices so this is double-checked right
/// before the call.
fn validate_transcode_buffer(
&mut self,
opts: &Options,
ptr: &TempLocal,
len: &TempLocal,
enc: FE,
) {
let mem_opts = match &opts.data_model {
DataModel::Gc {} => todo!("CM+GC"),
DataModel::LinearMemory(opts) => opts,
};
let width = enc.width();
self.validate_guest_pointer_align(mem_opts, ptr, width.into());

// `len` may come from the other memory, so its checks are done in 64
// bits.
let len_to_64 = |me: &mut Self| {
me.instruction(LocalGet(len.idx));
if len.ty == ValType::I32 {
me.instruction(I64ExtendI32U);
}
};

// The host doesn't access empty buffers, so there's nothing to load.
self.instruction(Block(BlockType::Empty));
len_to_64(self);
self.instruction(I64Eqz);
self.instruction(BrIf(0));

// Limiting the length to the maximum string size keeps the offset of
// the last code unit below 2 GiB, so it fits in a 32-bit pointer.
len_to_64(self);
let max = MAX_STRING_BYTE_LENGTH / u32::from(width);
self.instruction(I64Const(max.into()));
self.instruction(I64GtU);
self.instruction(If(BlockType::Empty));
self.trap(Trap::StringOutOfBounds);
self.instruction(End);

// Calculate the address of the last code unit, trapping if adding the
// offset to `ptr` wraps around.
self.instruction(LocalGet(ptr.idx));
len_to_64(self);
self.instruction(I64Const(1));
self.instruction(I64Sub);
if width == 2 {
self.instruction(I64Const(1));
self.instruction(I64Shl);
}
if !mem_opts.memory64() {
self.instruction(I32WrapI64);
}
self.ptr_add(mem_opts);
let last = self.local_tee_new_tmp(mem_opts.ptr());
self.instruction(LocalGet(ptr.idx));
self.ptr_lt_u(mem_opts);
self.instruction(If(BlockType::Empty));
self.trap(Trap::StringOutOfBounds);
self.instruction(End);

let last = Memory {
opts,
addr: last,
offset: 0,
};
match width {
1 => self.i32_load8u(&last),
_ => self.i32_load16u(&last),
}
self.instruction(Drop);
self.free_temp_local(last.addr);
self.instruction(End);
}

fn transcoder(
&mut self,
src: &WasmString<'_>,
Expand Down Expand Up @@ -3675,17 +3787,7 @@ impl<'a, 'b> Compiler<'a, 'b> {
DataModel::LinearMemory(mem_opts) => mem_opts,
};

// If the alignment is 1 then everything is trivially aligned and the
// check can be omitted.
if align != 1 {
self.instruction(LocalGet(addr.idx));
assert!(align.is_power_of_two());
self.ptr_uconst(mem_opts, align - 1);
self.ptr_and(mem_opts);
self.ptr_if(mem_opts, BlockType::Empty);
self.trap(Trap::UnalignedPointer);
self.instruction(End);
}
self.validate_guest_pointer_align(mem_opts, addr, align);

let extend_to_64 = |me: &mut Self| {
if !mem_opts.memory64() {
Expand Down Expand Up @@ -3735,6 +3837,27 @@ impl<'a, 'b> Compiler<'a, 'b> {
self.instruction(End);
}

/// Traps with `Trap::UnalignedPointer` unless `addr` is a multiple of
/// `align`.
fn validate_guest_pointer_align(
&mut self,
mem_opts: &LinearMemoryOptions,
addr: &TempLocal,
align: u32,
) {
// If the alignment is 1 then everything is trivially aligned and the
// check can be omitted.
if align != 1 {
self.instruction(LocalGet(addr.idx));
assert!(align.is_power_of_two());
self.ptr_uconst(mem_opts, align - 1);
self.ptr_and(mem_opts);
self.ptr_if(mem_opts, BlockType::Empty);
self.trap(Trap::UnalignedPointer);
self.instruction(End);
}
}

/// Generates a new local in this function of the `ty` specified,
/// initializing it with the top value on the current wasm stack.
///
Expand Down
Loading