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(any(not(feature = "worker-runtime"), feature = "native-runtime"), no_mangle)]
190pub unsafe extern "C" fn __fui_on_worker_progress(
193 worker_id: u32,
194 text_ptr: *const u8,
195 text_len: u32,
196) {
197 let message = if text_ptr.is_null() || text_len == 0 {
198 String::new()
199 } else {
200 String::from_utf8_lossy(unsafe { std::slice::from_raw_parts(text_ptr, text_len as usize) })
201 .into_owned()
202 };
203 with_active_worker(worker_id, |inner| {
204 if inner.finished || inner.cancel_requested {
205 return;
206 }
207 if let Some(callback) = inner.on_progress.clone() {
208 callback(WorkerProgressEventArgs { message });
209 }
210 });
211}
212
213#[cfg_attr(any(not(feature = "worker-runtime"), feature = "native-runtime"), no_mangle)]
214pub unsafe extern "C" fn __fui_on_worker_complete(
217 worker_id: u32,
218 text_ptr: *const u8,
219 text_len: u32,
220) {
221 let result = if text_ptr.is_null() || text_len == 0 {
222 String::new()
223 } else {
224 String::from_utf8_lossy(unsafe { std::slice::from_raw_parts(text_ptr, text_len as usize) })
225 .into_owned()
226 };
227 let Some(worker) = finish_worker(worker_id) else {
228 return;
229 };
230 let callback = {
231 let mut inner = worker.borrow_mut();
232 if inner.finished {
233 return;
234 }
235 inner.finished = true;
236 inner.on_complete.clone()
237 };
238 if let Some(callback) = callback {
239 callback(WorkerCompletedEventArgs { result });
240 }
241}
242
243#[cfg_attr(any(not(feature = "worker-runtime"), feature = "native-runtime"), no_mangle)]
244pub unsafe extern "C" fn __fui_on_worker_error(worker_id: u32, text_ptr: *const u8, text_len: u32) {
247 let message = if text_ptr.is_null() || text_len == 0 {
248 String::new()
249 } else {
250 String::from_utf8_lossy(unsafe { std::slice::from_raw_parts(text_ptr, text_len as usize) })
251 .into_owned()
252 };
253 let Some(worker) = finish_worker(worker_id) else {
254 return;
255 };
256 let callback = {
257 let mut inner = worker.borrow_mut();
258 if inner.finished {
259 return;
260 }
261 inner.finished = true;
262 inner.on_error.clone()
263 };
264 if let Some(callback) = callback {
265 callback(WorkerErrorEventArgs { message });
266 }
267}
268
269#[cfg(test)]
270mod tests {
271 use super::Worker;
272 use crate::ffi::{self, Call};
273 use std::cell::RefCell;
274 use std::rc::Rc;
275
276 fn started_worker_id(calls: &[Call]) -> u32 {
277 calls
278 .iter()
279 .find_map(|call| match call {
280 Call::WorkerStartString { worker_id, .. } => Some(*worker_id),
281 _ => None,
282 })
283 .expect("worker start call")
284 }
285
286 #[test]
287 fn worker_start_emits_host_call() {
288 ffi::test::reset();
289 let _worker = Worker::new("./workers/test.wasm", "demo").start("hello");
290 let calls = ffi::test::take_calls();
291 assert!(calls.iter().any(|call| matches!(call, Call::WorkerStartString { wasm_path, entry, input, .. } if wasm_path == "./workers/test.wasm" && entry == "demo" && input == "hello")));
292 }
293
294 #[test]
295 fn oversized_worker_start_input_reports_error_without_host_call() {
296 ffi::test::reset();
297 let error = Rc::new(RefCell::new(String::new()));
298 let error_clone = error.clone();
299 let input = "x".repeat(super::MAX_WORKER_START_INPUT_BYTES + 1);
300 let _worker = Worker::new("./workers/test.wasm", "demo")
301 .on_error(move |event| {
302 error_clone.replace(event.message);
303 })
304 .start(input);
305 let calls = ffi::test::take_calls();
306 assert!(!calls
307 .iter()
308 .any(|call| matches!(call, Call::WorkerStartString { .. })));
309 assert!(error.borrow().contains("maximum UTF-8 payload size"));
310 }
311
312 #[test]
313 fn worker_callbacks_receive_payloads() {
314 ffi::test::reset();
315 let progress = Rc::new(RefCell::new(String::new()));
316 let result = Rc::new(RefCell::new(String::new()));
317 let progress_clone = progress.clone();
318 let result_clone = result.clone();
319 let _worker = Worker::new("./workers/test.wasm", "demo")
320 .on_progress(move |event| {
321 progress_clone.replace(event.message);
322 })
323 .on_complete(move |event| {
324 result_clone.replace(event.result);
325 })
326 .start("hello");
327 let worker_id = started_worker_id(&ffi::test::take_calls());
328 unsafe {
329 super::__fui_on_worker_progress(worker_id, b"25%".as_ptr(), 3);
330 super::__fui_on_worker_complete(worker_id, b"done".as_ptr(), 4);
331 }
332 assert_eq!(&*progress.borrow(), "25%");
333 assert_eq!(&*result.borrow(), "done");
334 }
335
336 #[test]
337 fn worker_is_one_shot_and_first_terminal_callback_wins() {
338 ffi::test::reset();
339 let events = Rc::new(RefCell::new(Vec::new()));
340 let progress_events = events.clone();
341 let complete_events = events.clone();
342 let error_events = events.clone();
343 let worker = Worker::new("./workers/test.wasm", "demo")
344 .on_progress(move |event| progress_events.borrow_mut().push(event.message))
345 .on_complete(move |event| complete_events.borrow_mut().push(event.result))
346 .on_error(move |event| error_events.borrow_mut().push(event.message))
347 .start("first")
348 .start("second");
349 let calls = ffi::test::take_calls();
350 assert_eq!(
351 calls
352 .iter()
353 .filter(|call| matches!(call, Call::WorkerStartString { .. }))
354 .count(),
355 1
356 );
357 let worker_id = started_worker_id(&calls);
358 unsafe {
359 super::__fui_on_worker_progress(worker_id, b"progress".as_ptr(), 8);
360 super::__fui_on_worker_complete(worker_id, b"complete".as_ptr(), 8);
361 super::__fui_on_worker_error(worker_id, b"late error".as_ptr(), 10);
362 super::__fui_on_worker_progress(worker_id, b"late".as_ptr(), 4);
363 }
364 assert_eq!(&*events.borrow(), &["progress", "complete"]);
365 drop(worker);
366 }
367
368 #[test]
369 fn cancellation_is_idempotent_and_suppresses_progress() {
370 ffi::test::reset();
371 let progress = Rc::new(RefCell::new(Vec::new()));
372 let progress_clone = progress.clone();
373 let worker = Worker::new("./workers/test.wasm", "demo")
374 .on_progress(move |event| progress_clone.borrow_mut().push(event.message))
375 .start("input");
376 let start_calls = ffi::test::take_calls();
377 let worker_id = started_worker_id(&start_calls);
378 worker.cancel();
379 worker.cancel();
380 unsafe {
381 super::__fui_on_worker_progress(worker_id, b"ignored".as_ptr(), 7);
382 }
383 let calls = ffi::test::take_calls();
384 assert_eq!(
385 calls
386 .iter()
387 .filter(|call| matches!(call, Call::WorkerCancel { worker_id: id } if *id == worker_id))
388 .count(),
389 1
390 );
391 assert!(progress.borrow().is_empty());
392 }
393
394 #[test]
395 fn dropping_started_worker_cancels_and_detaches_callbacks() {
396 ffi::test::reset();
397 let completed = Rc::new(RefCell::new(false));
398 let completed_clone = completed.clone();
399 let worker = Worker::new("./workers/test.wasm", "demo")
400 .on_complete(move |_| {
401 completed_clone.replace(true);
402 })
403 .start("input");
404 let worker_id = started_worker_id(&ffi::test::take_calls());
405 drop(worker);
406 assert!(ffi::test::take_calls()
407 .iter()
408 .any(|call| matches!(call, Call::WorkerCancel { worker_id: id } if *id == worker_id)));
409 unsafe {
410 super::__fui_on_worker_complete(worker_id, b"late".as_ptr(), 4);
411 }
412 assert!(!*completed.borrow());
413 }
414}