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)]
193pub 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)]
220pub 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)]
253pub 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}