Skip to main content

fui/
worker.rs

1use std::cell::RefCell;
2use std::collections::HashMap;
3use std::rc::Rc;
4
5use crate::ffi;
6
7type ProgressCallback = Rc<dyn Fn(WorkerProgressEventArgs)>;
8type CompleteCallback = Rc<dyn Fn(WorkerCompletedEventArgs)>;
9type ErrorCallback = Rc<dyn Fn(WorkerErrorEventArgs)>;
10const MAX_WORKER_START_INPUT_BYTES: usize = 1024 * 1024;
11
12thread_local! {
13    static NEXT_WORKER_ID: RefCell<u32> = const { RefCell::new(1) };
14    static ACTIVE_WORKERS: RefCell<HashMap<u32, Rc<RefCell<WorkerInner>>>> = RefCell::new(HashMap::new());
15}
16
17fn with_utf8(value: &str, callback: impl FnOnce(usize, u32)) {
18    let bytes = value.as_bytes();
19    callback(
20        if bytes.is_empty() {
21            0
22        } else {
23            bytes.as_ptr() as usize
24        },
25        bytes.len() as u32,
26    );
27}
28
29#[derive(Clone, Debug, PartialEq, Eq)]
30pub struct WorkerProgressEventArgs {
31    pub message: String,
32}
33
34#[derive(Clone, Debug, PartialEq, Eq)]
35pub struct WorkerCompletedEventArgs {
36    pub result: String,
37}
38
39#[derive(Clone, Debug, PartialEq, Eq)]
40pub struct WorkerErrorEventArgs {
41    pub message: String,
42}
43
44struct WorkerInner {
45    worker_id: u32,
46    wasm_path: String,
47    entry_name: String,
48    on_progress: Option<ProgressCallback>,
49    on_complete: Option<CompleteCallback>,
50    on_error: Option<ErrorCallback>,
51    started: bool,
52    finished: bool,
53    cancel_requested: bool,
54}
55
56pub struct Worker {
57    inner: Rc<RefCell<WorkerInner>>,
58}
59
60impl Worker {
61    pub fn new(wasm_path: impl Into<String>, entry_name: impl Into<String>) -> Self {
62        let worker_id = NEXT_WORKER_ID.with(|next| {
63            let mut slot = next.borrow_mut();
64            let id = *slot;
65            *slot += 1;
66            id
67        });
68        let worker = Self {
69            inner: Rc::new(RefCell::new(WorkerInner {
70                worker_id,
71                wasm_path: wasm_path.into(),
72                entry_name: entry_name.into(),
73                on_progress: None,
74                on_complete: None,
75                on_error: None,
76                started: false,
77                finished: false,
78                cancel_requested: false,
79            })),
80        };
81        ACTIVE_WORKERS.with(|workers| {
82            workers.borrow_mut().insert(worker_id, worker.inner.clone());
83        });
84        worker
85    }
86
87    pub fn on_progress(self, handler: impl Fn(WorkerProgressEventArgs) + 'static) -> Self {
88        self.inner.borrow_mut().on_progress = Some(Rc::new(handler));
89        self
90    }
91
92    pub fn on_complete(self, handler: impl Fn(WorkerCompletedEventArgs) + 'static) -> Self {
93        self.inner.borrow_mut().on_complete = Some(Rc::new(handler));
94        self
95    }
96
97    pub fn on_error(self, handler: impl Fn(WorkerErrorEventArgs) + 'static) -> Self {
98        self.inner.borrow_mut().on_error = Some(Rc::new(handler));
99        self
100    }
101
102    pub fn start(self, input: impl Into<String>) -> Self {
103        let input = input.into();
104        let already_started = {
105            let inner = self.inner.borrow();
106            inner.started || inner.finished
107        };
108        if already_started {
109            return self;
110        }
111        if input.len() > MAX_WORKER_START_INPUT_BYTES {
112            let (worker_id, callback) = {
113                let mut inner = self.inner.borrow_mut();
114                inner.started = true;
115                inner.finished = true;
116                (inner.worker_id, inner.on_error.clone())
117            };
118            finish_worker(worker_id);
119            if let Some(callback) = callback {
120                callback(WorkerErrorEventArgs {
121                    message: format!(
122                        "Worker.start input exceeds the maximum UTF-8 payload size of {} bytes.",
123                        MAX_WORKER_START_INPUT_BYTES
124                    ),
125                });
126            }
127            return self;
128        }
129        let start_info = {
130            let mut inner = self.inner.borrow_mut();
131            inner.started = true;
132            (
133                inner.worker_id,
134                inner.wasm_path.clone(),
135                inner.entry_name.clone(),
136            )
137        };
138        with_utf8(&start_info.1, |wasm_path_ptr, wasm_path_len| {
139            with_utf8(&start_info.2, |entry_ptr, entry_len| {
140                with_utf8(&input, |input_ptr, input_len| unsafe {
141                    ffi::fui_worker_start_string(
142                        start_info.0,
143                        wasm_path_ptr,
144                        wasm_path_len,
145                        entry_ptr,
146                        entry_len,
147                        input_ptr,
148                        input_len,
149                    );
150                })
151            })
152        });
153        self
154    }
155
156    pub fn cancel(&self) {
157        let worker_id = {
158            let mut inner = self.inner.borrow_mut();
159            if !inner.started || inner.finished || inner.cancel_requested {
160                return;
161            }
162            inner.cancel_requested = true;
163            inner.worker_id
164        };
165        unsafe { ffi::fui_worker_cancel(worker_id) };
166    }
167}
168
169impl Drop for Worker {
170    fn drop(&mut self) {
171        self.cancel();
172        finish_worker(self.inner.borrow().worker_id);
173    }
174}
175
176fn finish_worker(worker_id: u32) -> Option<Rc<RefCell<WorkerInner>>> {
177    ACTIVE_WORKERS.with(|workers| workers.borrow_mut().remove(&worker_id))
178}
179
180fn with_active_worker(worker_id: u32, callback: impl FnOnce(&mut WorkerInner)) {
181    let Some(worker) = ACTIVE_WORKERS.with(|workers| workers.borrow().get(&worker_id).cloned())
182    else {
183        return;
184    };
185    let mut inner = worker.borrow_mut();
186    callback(&mut inner);
187}
188
189#[cfg_attr(
190    any(not(feature = "worker-runtime"), feature = "native-runtime"),
191    no_mangle
192)]
193/// # Safety
194/// `text_ptr` must be null for an empty message or point to `text_len` readable bytes.
195pub unsafe extern "C" fn __fui_on_worker_progress(
196    worker_id: u32,
197    text_ptr: *const u8,
198    text_len: u32,
199) {
200    let message = if text_ptr.is_null() || text_len == 0 {
201        String::new()
202    } else {
203        String::from_utf8_lossy(unsafe { std::slice::from_raw_parts(text_ptr, text_len as usize) })
204            .into_owned()
205    };
206    with_active_worker(worker_id, |inner| {
207        if inner.finished || inner.cancel_requested {
208            return;
209        }
210        if let Some(callback) = inner.on_progress.clone() {
211            callback(WorkerProgressEventArgs { message });
212        }
213    });
214}
215
216#[cfg_attr(
217    any(not(feature = "worker-runtime"), feature = "native-runtime"),
218    no_mangle
219)]
220/// # Safety
221/// `text_ptr` must be null for an empty result or point to `text_len` readable bytes.
222pub unsafe extern "C" fn __fui_on_worker_complete(
223    worker_id: u32,
224    text_ptr: *const u8,
225    text_len: u32,
226) {
227    let result = if text_ptr.is_null() || text_len == 0 {
228        String::new()
229    } else {
230        String::from_utf8_lossy(unsafe { std::slice::from_raw_parts(text_ptr, text_len as usize) })
231            .into_owned()
232    };
233    let Some(worker) = finish_worker(worker_id) else {
234        return;
235    };
236    let callback = {
237        let mut inner = worker.borrow_mut();
238        if inner.finished {
239            return;
240        }
241        inner.finished = true;
242        inner.on_complete.clone()
243    };
244    if let Some(callback) = callback {
245        callback(WorkerCompletedEventArgs { result });
246    }
247}
248
249#[cfg_attr(
250    any(not(feature = "worker-runtime"), feature = "native-runtime"),
251    no_mangle
252)]
253/// # Safety
254/// `text_ptr` must be null for an empty message or point to `text_len` readable bytes.
255pub unsafe extern "C" fn __fui_on_worker_error(worker_id: u32, text_ptr: *const u8, text_len: u32) {
256    let message = if text_ptr.is_null() || text_len == 0 {
257        String::new()
258    } else {
259        String::from_utf8_lossy(unsafe { std::slice::from_raw_parts(text_ptr, text_len as usize) })
260            .into_owned()
261    };
262    let Some(worker) = finish_worker(worker_id) else {
263        return;
264    };
265    let callback = {
266        let mut inner = worker.borrow_mut();
267        if inner.finished {
268            return;
269        }
270        inner.finished = true;
271        inner.on_error.clone()
272    };
273    if let Some(callback) = callback {
274        callback(WorkerErrorEventArgs { message });
275    }
276}
277
278#[cfg(test)]
279mod tests {
280    use super::Worker;
281    use crate::ffi::{self, Call};
282    use std::cell::RefCell;
283    use std::rc::Rc;
284
285    fn started_worker_id(calls: &[Call]) -> u32 {
286        calls
287            .iter()
288            .find_map(|call| match call {
289                Call::WorkerStartString { worker_id, .. } => Some(*worker_id),
290                _ => None,
291            })
292            .expect("worker start call")
293    }
294
295    #[test]
296    fn worker_start_emits_host_call() {
297        ffi::test::reset();
298        let _worker = Worker::new("./workers/test.wasm", "demo").start("hello");
299        let calls = ffi::test::take_calls();
300        assert!(calls.iter().any(|call| matches!(call, Call::WorkerStartString { wasm_path, entry, input, .. } if wasm_path == "./workers/test.wasm" && entry == "demo" && input == "hello")));
301    }
302
303    #[test]
304    fn oversized_worker_start_input_reports_error_without_host_call() {
305        ffi::test::reset();
306        let error = Rc::new(RefCell::new(String::new()));
307        let error_clone = error.clone();
308        let input = "x".repeat(super::MAX_WORKER_START_INPUT_BYTES + 1);
309        let _worker = Worker::new("./workers/test.wasm", "demo")
310            .on_error(move |event| {
311                error_clone.replace(event.message);
312            })
313            .start(input);
314        let calls = ffi::test::take_calls();
315        assert!(!calls
316            .iter()
317            .any(|call| matches!(call, Call::WorkerStartString { .. })));
318        assert!(error.borrow().contains("maximum UTF-8 payload size"));
319    }
320
321    #[test]
322    fn worker_callbacks_receive_payloads() {
323        ffi::test::reset();
324        let progress = Rc::new(RefCell::new(String::new()));
325        let result = Rc::new(RefCell::new(String::new()));
326        let progress_clone = progress.clone();
327        let result_clone = result.clone();
328        let _worker = Worker::new("./workers/test.wasm", "demo")
329            .on_progress(move |event| {
330                progress_clone.replace(event.message);
331            })
332            .on_complete(move |event| {
333                result_clone.replace(event.result);
334            })
335            .start("hello");
336        let worker_id = started_worker_id(&ffi::test::take_calls());
337        unsafe {
338            super::__fui_on_worker_progress(worker_id, b"25%".as_ptr(), 3);
339            super::__fui_on_worker_complete(worker_id, b"done".as_ptr(), 4);
340        }
341        assert_eq!(&*progress.borrow(), "25%");
342        assert_eq!(&*result.borrow(), "done");
343    }
344
345    #[test]
346    fn worker_is_one_shot_and_first_terminal_callback_wins() {
347        ffi::test::reset();
348        let events = Rc::new(RefCell::new(Vec::new()));
349        let progress_events = events.clone();
350        let complete_events = events.clone();
351        let error_events = events.clone();
352        let worker = Worker::new("./workers/test.wasm", "demo")
353            .on_progress(move |event| progress_events.borrow_mut().push(event.message))
354            .on_complete(move |event| complete_events.borrow_mut().push(event.result))
355            .on_error(move |event| error_events.borrow_mut().push(event.message))
356            .start("first")
357            .start("second");
358        let calls = ffi::test::take_calls();
359        assert_eq!(
360            calls
361                .iter()
362                .filter(|call| matches!(call, Call::WorkerStartString { .. }))
363                .count(),
364            1
365        );
366        let worker_id = started_worker_id(&calls);
367        unsafe {
368            super::__fui_on_worker_progress(worker_id, b"progress".as_ptr(), 8);
369            super::__fui_on_worker_complete(worker_id, b"complete".as_ptr(), 8);
370            super::__fui_on_worker_error(worker_id, b"late error".as_ptr(), 10);
371            super::__fui_on_worker_progress(worker_id, b"late".as_ptr(), 4);
372        }
373        assert_eq!(&*events.borrow(), &["progress", "complete"]);
374        drop(worker);
375    }
376
377    #[test]
378    fn cancellation_is_idempotent_and_suppresses_progress() {
379        ffi::test::reset();
380        let progress = Rc::new(RefCell::new(Vec::new()));
381        let progress_clone = progress.clone();
382        let worker = Worker::new("./workers/test.wasm", "demo")
383            .on_progress(move |event| progress_clone.borrow_mut().push(event.message))
384            .start("input");
385        let start_calls = ffi::test::take_calls();
386        let worker_id = started_worker_id(&start_calls);
387        worker.cancel();
388        worker.cancel();
389        unsafe {
390            super::__fui_on_worker_progress(worker_id, b"ignored".as_ptr(), 7);
391        }
392        let calls = ffi::test::take_calls();
393        assert_eq!(
394            calls
395                .iter()
396                .filter(
397                    |call| matches!(call, Call::WorkerCancel { worker_id: id } if *id == worker_id)
398                )
399                .count(),
400            1
401        );
402        assert!(progress.borrow().is_empty());
403    }
404
405    #[test]
406    fn dropping_started_worker_cancels_and_detaches_callbacks() {
407        ffi::test::reset();
408        let completed = Rc::new(RefCell::new(false));
409        let completed_clone = completed.clone();
410        let worker = Worker::new("./workers/test.wasm", "demo")
411            .on_complete(move |_| {
412                completed_clone.replace(true);
413            })
414            .start("input");
415        let worker_id = started_worker_id(&ffi::test::take_calls());
416        drop(worker);
417        assert!(ffi::test::take_calls()
418            .iter()
419            .any(|call| matches!(call, Call::WorkerCancel { worker_id: id } if *id == worker_id)));
420        unsafe {
421            super::__fui_on_worker_complete(worker_id, b"late".as_ptr(), 4);
422        }
423        assert!(!*completed.borrow());
424    }
425}