Skip to main content

deterministic_wasi_ctx/
scheduling.rs

1use std::{mem, ptr, slice};
2
3use wasi::{Event, EventFdReadwrite, Subscription};
4use wasmtime::{format_err, Caller, Linker, Result};
5
6/// Adds implementations for WASI preview 1 `poll_oneoff` and `sched_yield` to
7/// the linker which will return immediately.
8/// Note: This function will enable shadowing on the linker.
9pub fn replace_scheduling_functions<T>(linker: &mut Linker<T>) -> Result<()>
10where
11    T: Send + 'static,
12{
13    override_scheduling_functions(linker, "wasi_snapshot_preview1")
14}
15
16/// Adds implementations for WASI preview 0 `poll_oneoff` and `sched_yield` to
17/// the linker which will return immediately.
18/// Note: This function will enable shadowing on the linker.
19pub fn replace_scheduling_functions_for_wasi_preview_0<T>(linker: &mut Linker<T>) -> Result<()>
20where
21    T: Send + 'static,
22{
23    override_scheduling_functions(linker, "wasi_unstable")
24}
25
26fn override_scheduling_functions<T: 'static>(linker: &mut Linker<T>, module: &str) -> Result<()> {
27    linker.allow_shadowing(true);
28    linker.func_wrap(
29        module,
30        "poll_oneoff",
31        |mut caller: Caller<'_, T>,
32         in_ptr: i32,
33         out_ptr: i32,
34         nsubscriptions: i32,
35         nevents_ptr: i32|
36         -> Result<i32> {
37            let in_ptr = in_ptr as usize;
38            let out_ptr = out_ptr as usize;
39            let nsubscriptions = nsubscriptions as usize;
40            let nevents_ptr = nevents_ptr as usize;
41            // See https://github.com/WebAssembly/WASI/blob/3d5e0553cd01dd4d6e2c06ad2a702ee9dda17b7f/legacy/tools/witx-docs.md#pointers
42            let memory = caller
43                .get_export("memory")
44                .map_or_else(|| Err(format_err!("missing required memory export")), Ok)?
45                .into_memory()
46                .map_or_else(|| Err(format_err!("missing required memory export")), Ok)?;
47
48            for i in 0..nsubscriptions {
49                // Read the `Subscription` from memory from `in_ptr`.
50                let offset = in_ptr + (i * mem::size_of::<Subscription>());
51                let mut subscription_buffer = [0u8; mem::size_of::<Subscription>()];
52                memory.read(&caller, offset, &mut subscription_buffer)?;
53                let subscription =
54                    unsafe { ptr::read(subscription_buffer.as_ptr() as *const Subscription) };
55
56                // Create a successful `Event` for each subscription.
57                let event = Event {
58                    userdata: subscription.userdata,
59                    error: wasi::ERRNO_SUCCESS,
60                    // See https://github.com/WebAssembly/wasi-libc/blob/e9524a0980b9bb6bb92e87a41ed1055bdda5bb86/libc-bottom-half/headers/public/wasi/api.h#L1100-L1121
61                    // for the mapping between the integers and the event type.
62                    type_: match subscription.u.tag {
63                        0 => wasi::EVENTTYPE_CLOCK,
64                        1 => wasi::EVENTTYPE_FD_READ,
65                        2 => wasi::EVENTTYPE_FD_WRITE,
66                        _ => unreachable!(),
67                    },
68                    fd_readwrite: EventFdReadwrite {
69                        nbytes: 0,
70                        flags: 0,
71                    },
72                };
73
74                // Write the event into memory at `out_ptr`.
75                let offset = out_ptr + (i * mem::size_of::<Event>());
76                let event_buffer = unsafe {
77                    slice::from_raw_parts(
78                        &event as *const Event as *const u8,
79                        mem::size_of::<Event>(),
80                    )
81                };
82                memory.write(&mut caller, offset, event_buffer)?
83            }
84
85            // Copy number of subscriptions into number of events pointer.
86            let buffer = nsubscriptions.to_le_bytes();
87            memory.write(&mut caller, nevents_ptr, &buffer)?;
88
89            Ok(wasi::ERRNO_SUCCESS.raw() as i32)
90        },
91    )?;
92
93    linker.func_wrap(module, "sched_yield", || wasi::ERRNO_SUCCESS.raw() as i32)?;
94
95    Ok(())
96}