deterministic_wasi_ctx/
scheduling.rs

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