1 use crate::component::func::{LiftContext, LowerContext, Options};
2 use crate::component::matching::InstanceType;
3 use crate::component::storage::slice_to_storage_mut;
4 use crate::component::{ComponentNamedList, ComponentType, Lift, Lower, Val};
5 use crate::prelude::*;
6 use crate::runtime::vm::component::{
7     ComponentInstance, InstanceFlags, VMComponentContext, VMLowering, VMLoweringCallee,
8 };
9 use crate::runtime::vm::{VMFuncRef, VMGlobalDefinition, VMMemoryDefinition, VMOpaqueContext};
10 use crate::{AsContextMut, CallHook, StoreContextMut, ValRaw};
11 use alloc::sync::Arc;
12 use core::any::Any;
13 use core::mem::{self, MaybeUninit};
14 use core::ptr::NonNull;
15 use wasmtime_environ::component::{
16     CanonicalAbiInfo, ComponentTypes, InterfaceType, MAX_FLAT_PARAMS, MAX_FLAT_RESULTS,
17     StringEncoding, TypeFuncIndex,
18 };
19 
20 pub struct HostFunc {
21     entrypoint: VMLoweringCallee,
22     typecheck: Box<dyn (Fn(TypeFuncIndex, &InstanceType<'_>) -> Result<()>) + Send + Sync>,
23     func: Box<dyn Any + Send + Sync>,
24 }
25 
26 impl core::fmt::Debug for HostFunc {
27     fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
28         f.debug_struct("HostFunc").finish_non_exhaustive()
29     }
30 }
31 
32 impl HostFunc {
33     pub(crate) fn from_closure<T, F, P, R>(func: F) -> Arc<HostFunc>
34     where
35         F: Fn(StoreContextMut<T>, P) -> Result<R> + Send + Sync + 'static,
36         P: ComponentNamedList + Lift + 'static,
37         R: ComponentNamedList + Lower + 'static,
38         T: 'static,
39     {
40         let entrypoint = Self::entrypoint::<T, F, P, R>;
41         Arc::new(HostFunc {
42             entrypoint,
43             typecheck: Box::new(typecheck::<P, R>),
44             func: Box::new(func),
45         })
46     }
47 
48     extern "C" fn entrypoint<T, F, P, R>(
49         cx: NonNull<VMOpaqueContext>,
50         data: NonNull<u8>,
51         ty: u32,
52         _caller_instance: u32,
53         flags: NonNull<VMGlobalDefinition>,
54         memory: *mut VMMemoryDefinition,
55         realloc: *mut VMFuncRef,
56         string_encoding: u8,
57         async_: u8,
58         storage: NonNull<MaybeUninit<ValRaw>>,
59         storage_len: usize,
60     ) -> bool
61     where
62         F: Fn(StoreContextMut<T>, P) -> Result<R>,
63         P: ComponentNamedList + Lift + 'static,
64         R: ComponentNamedList + Lower + 'static,
65         T: 'static,
66     {
67         let data = data.as_ptr() as *const F;
68         unsafe {
69             call_host_and_handle_result::<T>(cx, |instance, types, store| {
70                 call_host::<_, _, _, _>(
71                     instance,
72                     types,
73                     store,
74                     TypeFuncIndex::from_u32(ty),
75                     InstanceFlags::from_raw(flags),
76                     memory,
77                     realloc,
78                     StringEncoding::from_u8(string_encoding).unwrap(),
79                     async_ != 0,
80                     NonNull::slice_from_raw_parts(storage, storage_len).as_mut(),
81                     |store, args| (*data)(store, args),
82                 )
83             })
84         }
85     }
86 
87     pub(crate) fn new_dynamic<T, F>(func: F) -> Arc<HostFunc>
88     where
89         F: Fn(StoreContextMut<'_, T>, &[Val], &mut [Val]) -> Result<()> + Send + Sync + 'static,
90         T: 'static,
91     {
92         Arc::new(HostFunc {
93             entrypoint: dynamic_entrypoint::<T, F>,
94             // This function performs dynamic type checks and subsequently does
95             // not need to perform up-front type checks. Instead everything is
96             // dynamically managed at runtime.
97             typecheck: Box::new(move |_expected_index, _expected_types| Ok(())),
98             func: Box::new(func),
99         })
100     }
101 
102     pub fn typecheck(&self, ty: TypeFuncIndex, types: &InstanceType<'_>) -> Result<()> {
103         (self.typecheck)(ty, types)
104     }
105 
106     pub fn lowering(&self) -> VMLowering {
107         let data = NonNull::from(&*self.func).cast();
108         VMLowering {
109             callee: self.entrypoint,
110             data: data.into(),
111         }
112     }
113 }
114 
115 fn typecheck<P, R>(ty: TypeFuncIndex, types: &InstanceType<'_>) -> Result<()>
116 where
117     P: ComponentNamedList + Lift,
118     R: ComponentNamedList + Lower,
119 {
120     let ty = &types.types[ty];
121     P::typecheck(&InterfaceType::Tuple(ty.params), types)
122         .context("type mismatch with parameters")?;
123     R::typecheck(&InterfaceType::Tuple(ty.results), types).context("type mismatch with results")?;
124     Ok(())
125 }
126 
127 /// The "meat" of calling a host function from wasm.
128 ///
129 /// This function is delegated to from implementations of
130 /// `HostFunc::from_closure`. Most of the arguments from the `entrypoint` are
131 /// forwarded here except for the `data` pointer which is encapsulated in the
132 /// `closure` argument here.
133 ///
134 /// This function is parameterized over:
135 ///
136 /// * `T` - the type of store this function works with (an unsafe assertion)
137 /// * `Params` - the parameters to the host function, viewed as a tuple
138 /// * `Return` - the result of the host function
139 /// * `F` - the `closure` to actually receive the `Params` and return the
140 ///   `Return`
141 ///
142 /// It's expected that `F` will "un-tuple" the arguments to pass to a host
143 /// closure.
144 ///
145 /// This function is in general `unsafe` as the validity of all the parameters
146 /// must be upheld. Generally that's done by ensuring this is only called from
147 /// the select few places it's intended to be called from.
148 unsafe fn call_host<T, Params, Return, F>(
149     instance: *mut ComponentInstance,
150     types: &Arc<ComponentTypes>,
151     mut cx: StoreContextMut<'_, T>,
152     ty: TypeFuncIndex,
153     mut flags: InstanceFlags,
154     memory: *mut VMMemoryDefinition,
155     realloc: *mut VMFuncRef,
156     string_encoding: StringEncoding,
157     async_: bool,
158     storage: &mut [MaybeUninit<ValRaw>],
159     closure: F,
160 ) -> Result<()>
161 where
162     Params: Lift,
163     Return: Lower,
164     F: FnOnce(StoreContextMut<'_, T>, Params) -> Result<Return>,
165 {
166     if async_ {
167         todo!()
168     }
169 
170     /// Representation of arguments to this function when a return pointer is in
171     /// use, namely the argument list is followed by a single value which is the
172     /// return pointer.
173     #[repr(C)]
174     struct ReturnPointer<T> {
175         args: T,
176         retptr: ValRaw,
177     }
178 
179     /// Representation of arguments to this function when the return value is
180     /// returned directly, namely the arguments and return value all start from
181     /// the beginning (aka this is a `union`, not a `struct`).
182     #[repr(C)]
183     union ReturnStack<T: Copy, U: Copy> {
184         args: T,
185         ret: U,
186     }
187 
188     let options = Options::new(
189         cx.0.id(),
190         NonNull::new(memory),
191         NonNull::new(realloc),
192         string_encoding,
193     );
194 
195     // Perform a dynamic check that this instance can indeed be left. Exiting
196     // the component is disallowed, for example, when the `realloc` function
197     // calls a canonical import.
198     if !flags.may_leave() {
199         bail!("cannot leave component instance");
200     }
201 
202     let ty = &types[ty];
203     let param_tys = InterfaceType::Tuple(ty.params);
204     let result_tys = InterfaceType::Tuple(ty.results);
205 
206     // There's a 2x2 matrix of whether parameters and results are stored on the
207     // stack or on the heap. Each of the 4 branches here have a different
208     // representation of the storage of arguments/returns.
209     //
210     // Also note that while four branches are listed here only one is taken for
211     // any particular `Params` and `Return` combination. This should be
212     // trivially DCE'd by LLVM. Perhaps one day with enough const programming in
213     // Rust we can make monomorphizations of this function codegen only one
214     // branch, but today is not that day.
215     let mut storage: Storage<'_, Params, Return> = if Params::flatten_count() <= MAX_FLAT_PARAMS {
216         if Return::flatten_count() <= MAX_FLAT_RESULTS {
217             Storage::Direct(slice_to_storage_mut(storage))
218         } else {
219             Storage::ResultsIndirect(slice_to_storage_mut(storage).assume_init_ref())
220         }
221     } else {
222         if Return::flatten_count() <= MAX_FLAT_RESULTS {
223             Storage::ParamsIndirect(slice_to_storage_mut(storage))
224         } else {
225             Storage::Indirect(slice_to_storage_mut(storage).assume_init_ref())
226         }
227     };
228     let mut lift = LiftContext::new(cx.0, &options, types, instance);
229     lift.enter_call();
230     let params = storage.lift_params(&mut lift, param_tys)?;
231 
232     let ret = closure(cx.as_context_mut(), params)?;
233     flags.set_may_leave(false);
234     let mut lower = LowerContext::new(cx, &options, types, instance);
235     storage.lower_results(&mut lower, result_tys, ret)?;
236     flags.set_may_leave(true);
237 
238     lower.exit_call()?;
239 
240     return Ok(());
241 
242     enum Storage<'a, P: ComponentType, R: ComponentType> {
243         Direct(&'a mut MaybeUninit<ReturnStack<P::Lower, R::Lower>>),
244         ParamsIndirect(&'a mut MaybeUninit<ReturnStack<ValRaw, R::Lower>>),
245         ResultsIndirect(&'a ReturnPointer<P::Lower>),
246         Indirect(&'a ReturnPointer<ValRaw>),
247     }
248 
249     impl<P, R> Storage<'_, P, R>
250     where
251         P: ComponentType + Lift,
252         R: ComponentType + Lower,
253     {
254         unsafe fn lift_params(&self, cx: &mut LiftContext<'_>, ty: InterfaceType) -> Result<P> {
255             match self {
256                 Storage::Direct(storage) => P::lift(cx, ty, &storage.assume_init_ref().args),
257                 Storage::ResultsIndirect(storage) => P::lift(cx, ty, &storage.args),
258                 Storage::ParamsIndirect(storage) => {
259                     let ptr = validate_inbounds::<P>(cx.memory(), &storage.assume_init_ref().args)?;
260                     P::load(cx, ty, &cx.memory()[ptr..][..P::SIZE32])
261                 }
262                 Storage::Indirect(storage) => {
263                     let ptr = validate_inbounds::<P>(cx.memory(), &storage.args)?;
264                     P::load(cx, ty, &cx.memory()[ptr..][..P::SIZE32])
265                 }
266             }
267         }
268 
269         unsafe fn lower_results<T>(
270             &mut self,
271             cx: &mut LowerContext<'_, T>,
272             ty: InterfaceType,
273             ret: R,
274         ) -> Result<()> {
275             match self {
276                 Storage::Direct(storage) => ret.lower(cx, ty, map_maybe_uninit!(storage.ret)),
277                 Storage::ParamsIndirect(storage) => {
278                     ret.lower(cx, ty, map_maybe_uninit!(storage.ret))
279                 }
280                 Storage::ResultsIndirect(storage) => {
281                     let ptr = validate_inbounds::<R>(cx.as_slice_mut(), &storage.retptr)?;
282                     ret.store(cx, ty, ptr)
283                 }
284                 Storage::Indirect(storage) => {
285                     let ptr = validate_inbounds::<R>(cx.as_slice_mut(), &storage.retptr)?;
286                     ret.store(cx, ty, ptr)
287                 }
288             }
289         }
290     }
291 }
292 
293 fn validate_inbounds<T: ComponentType>(memory: &[u8], ptr: &ValRaw) -> Result<usize> {
294     // FIXME(#4311): needs memory64 support
295     let ptr = usize::try_from(ptr.get_u32())?;
296     if ptr % usize::try_from(T::ALIGN32)? != 0 {
297         bail!("pointer not aligned");
298     }
299     let end = match ptr.checked_add(T::SIZE32) {
300         Some(n) => n,
301         None => bail!("pointer size overflow"),
302     };
303     if end > memory.len() {
304         bail!("pointer out of bounds")
305     }
306     Ok(ptr)
307 }
308 
309 unsafe fn call_host_and_handle_result<T>(
310     cx: NonNull<VMOpaqueContext>,
311     func: impl FnOnce(
312         *mut ComponentInstance,
313         &Arc<ComponentTypes>,
314         StoreContextMut<'_, T>,
315     ) -> Result<()>,
316 ) -> bool
317 where
318     T: 'static,
319 {
320     let cx = VMComponentContext::from_opaque(cx);
321     let instance = cx.as_ref().instance();
322     let types = (*instance).component_types();
323     let raw_store = (*instance).store();
324     let mut store = StoreContextMut(&mut *raw_store.cast());
325 
326     crate::runtime::vm::catch_unwind_and_record_trap(|| {
327         store.0.call_hook(CallHook::CallingHost)?;
328         let res = func(instance, types, store.as_context_mut());
329         store.0.call_hook(CallHook::ReturningFromHost)?;
330         res
331     })
332 }
333 
334 unsafe fn call_host_dynamic<T, F>(
335     instance: *mut ComponentInstance,
336     types: &Arc<ComponentTypes>,
337     mut store: StoreContextMut<'_, T>,
338     ty: TypeFuncIndex,
339     mut flags: InstanceFlags,
340     memory: *mut VMMemoryDefinition,
341     realloc: *mut VMFuncRef,
342     string_encoding: StringEncoding,
343     async_: bool,
344     storage: &mut [MaybeUninit<ValRaw>],
345     closure: F,
346 ) -> Result<()>
347 where
348     F: FnOnce(StoreContextMut<'_, T>, &[Val], &mut [Val]) -> Result<()>,
349     T: 'static,
350 {
351     if async_ {
352         todo!()
353     }
354 
355     let options = Options::new(
356         store.0.id(),
357         NonNull::new(memory),
358         NonNull::new(realloc),
359         string_encoding,
360     );
361 
362     // Perform a dynamic check that this instance can indeed be left. Exiting
363     // the component is disallowed, for example, when the `realloc` function
364     // calls a canonical import.
365     if !flags.may_leave() {
366         bail!("cannot leave component instance");
367     }
368 
369     let args;
370     let ret_index;
371 
372     let func_ty = &types[ty];
373     let param_tys = &types[func_ty.params];
374     let result_tys = &types[func_ty.results];
375     let mut cx = LiftContext::new(store.0, &options, types, instance);
376     cx.enter_call();
377     if let Some(param_count) = param_tys.abi.flat_count(MAX_FLAT_PARAMS) {
378         // NB: can use `MaybeUninit::slice_assume_init_ref` when that's stable
379         let mut iter =
380             mem::transmute::<&[MaybeUninit<ValRaw>], &[ValRaw]>(&storage[..param_count]).iter();
381         args = param_tys
382             .types
383             .iter()
384             .map(|ty| Val::lift(&mut cx, *ty, &mut iter))
385             .collect::<Result<Box<[_]>>>()?;
386         ret_index = param_count;
387         assert!(iter.next().is_none());
388     } else {
389         let mut offset =
390             validate_inbounds_dynamic(&param_tys.abi, cx.memory(), storage[0].assume_init_ref())?;
391         args = param_tys
392             .types
393             .iter()
394             .map(|ty| {
395                 let abi = types.canonical_abi(ty);
396                 let size = usize::try_from(abi.size32).unwrap();
397                 let memory = &cx.memory()[abi.next_field32_size(&mut offset)..][..size];
398                 Val::load(&mut cx, *ty, memory)
399             })
400             .collect::<Result<Box<[_]>>>()?;
401         ret_index = 1;
402     };
403 
404     let mut result_vals = Vec::with_capacity(result_tys.types.len());
405     for _ in result_tys.types.iter() {
406         result_vals.push(Val::Bool(false));
407     }
408     closure(store.as_context_mut(), &args, &mut result_vals)?;
409     flags.set_may_leave(false);
410 
411     let mut cx = LowerContext::new(store, &options, types, instance);
412     if let Some(cnt) = result_tys.abi.flat_count(MAX_FLAT_RESULTS) {
413         let mut dst = storage[..cnt].iter_mut();
414         for (val, ty) in result_vals.iter().zip(result_tys.types.iter()) {
415             val.lower(&mut cx, *ty, &mut dst)?;
416         }
417         assert!(dst.next().is_none());
418     } else {
419         let ret_ptr = storage[ret_index].assume_init_ref();
420         let mut ptr = validate_inbounds_dynamic(&result_tys.abi, cx.as_slice_mut(), ret_ptr)?;
421         for (val, ty) in result_vals.iter().zip(result_tys.types.iter()) {
422             let offset = types.canonical_abi(ty).next_field32_size(&mut ptr);
423             val.store(&mut cx, *ty, offset)?;
424         }
425     }
426 
427     flags.set_may_leave(true);
428 
429     cx.exit_call()?;
430 
431     return Ok(());
432 }
433 
434 fn validate_inbounds_dynamic(abi: &CanonicalAbiInfo, memory: &[u8], ptr: &ValRaw) -> Result<usize> {
435     // FIXME(#4311): needs memory64 support
436     let ptr = usize::try_from(ptr.get_u32())?;
437     if ptr % usize::try_from(abi.align32)? != 0 {
438         bail!("pointer not aligned");
439     }
440     let end = match ptr.checked_add(usize::try_from(abi.size32).unwrap()) {
441         Some(n) => n,
442         None => bail!("pointer size overflow"),
443     };
444     if end > memory.len() {
445         bail!("pointer out of bounds")
446     }
447     Ok(ptr)
448 }
449 
450 extern "C" fn dynamic_entrypoint<T, F>(
451     cx: NonNull<VMOpaqueContext>,
452     data: NonNull<u8>,
453     ty: u32,
454     _caller_instance: u32,
455     flags: NonNull<VMGlobalDefinition>,
456     memory: *mut VMMemoryDefinition,
457     realloc: *mut VMFuncRef,
458     string_encoding: u8,
459     async_: u8,
460     storage: NonNull<MaybeUninit<ValRaw>>,
461     storage_len: usize,
462 ) -> bool
463 where
464     F: Fn(StoreContextMut<'_, T>, &[Val], &mut [Val]) -> Result<()> + Send + Sync + 'static,
465     T: 'static,
466 {
467     let data = data.as_ptr() as *const F;
468     unsafe {
469         call_host_and_handle_result(cx, |instance, types, store| {
470             call_host_dynamic::<T, _>(
471                 instance,
472                 types,
473                 store,
474                 TypeFuncIndex::from_u32(ty),
475                 InstanceFlags::from_raw(flags),
476                 memory,
477                 realloc,
478                 StringEncoding::from_u8(string_encoding).unwrap(),
479                 async_ != 0,
480                 NonNull::slice_from_raw_parts(storage, storage_len).as_mut(),
481                 |store, params, results| (*data)(store, params, results),
482             )
483         })
484     }
485 }
486