1use std::cell::Cell;
10use std::fmt;
11use std::rc::Rc;
12use std::sync::Arc;
13
14use crate::ctx::Ctx;
15use crate::dispatcher::{PoolKind, WorkerSubmitError};
16use crate::graph::{Graph, GraphNodeOpts};
17use crate::node::{Core, Node, NodeOpts};
18use crate::operators::Operator;
19use crate::protocol::Message;
20
21pub fn worker_derived<I, R, E, P, C>(
29 graph: &Graph,
30 deps: Vec<Core>,
31 prepare: P,
32 compute: C,
33 mut opts: GraphNodeOpts,
34) -> Node<R>
35where
36 I: Send + 'static,
37 R: Send + 'static,
38 E: fmt::Display + Send + 'static,
39 P: Fn(&Ctx) -> Option<I> + 'static,
40 C: Fn(I) -> Result<R, E> + Send + Sync + 'static,
41{
42 opts.node.pool = PoolKind::Async;
43 let compute = Arc::new(compute);
44 let latest_invocation = Rc::new(Cell::new(0u64));
45 let op = Operator::with_opts(
46 "workerDerived",
47 NodeOpts {
48 complete_when_deps_complete: false,
51 pool: PoolKind::Async,
52 ..NodeOpts::default()
53 },
54 move |ctx| {
55 let latest_invocation = latest_invocation.clone();
56 let invocation = latest_invocation
57 .get()
58 .checked_add(1)
59 .expect("worker_derived invocation generation overflow");
60 latest_invocation.set(invocation);
61 let Some(input) = prepare(ctx) else {
62 ctx.down(vec![Message::Resolved]);
63 return;
64 };
65 let Some(driver) = ctx.local_async_driver() else {
66 ctx.down(vec![Message::Error(
67 "worker_derived: missing local async driver".into(),
68 )]);
69 return;
70 };
71 let out = ctx.defer();
72 let compute = compute.clone();
73 let job = match ctx.dispatcher().submit_worker(input, compute) {
74 Ok(job) => job,
75 Err(WorkerSubmitError::MissingBackend) => {
76 ctx.down(vec![Message::Error(
77 "worker_derived: missing worker backend".into(),
78 )]);
79 return;
80 }
81 Err(WorkerSubmitError::MissingRuntime) => {
82 ctx.down(vec![Message::Error(
83 "worker_derived: missing Tokio runtime".into(),
84 )]);
85 return;
86 }
87 };
88 let cancel = driver.spawn_local(Box::pin(async move {
89 let joined = job.spawn().await;
90 if latest_invocation.get() != invocation {
91 return;
92 }
93 match joined {
94 Ok(Ok(value)) => out.emit(value),
95 Ok(Err(error)) => out.down(vec![Message::Error(error.into())]),
96 Err(error) => out.down(vec![Message::Error(
97 format!("worker_derived: worker task failed: {error}").into(),
98 )]),
99 }
100 }));
101 ctx.on_deactivation(cancel);
102 },
103 );
104 graph.init_node(op, deps, opts)
105}
106
107#[cfg(test)]
108mod tests {
109 use super::*;
110 use std::cell::RefCell;
111 use std::future::Future;
112 use std::pin::Pin;
113 use std::rc::Rc;
114 use std::sync::atomic::{AtomicUsize, Ordering};
115 use std::sync::{Arc, Mutex};
116 use std::thread::ThreadId;
117 use std::time::Duration;
118
119 use crate::async_driver::{DriverCancel, LocalAsyncDriver, TokioLocalDriver};
120 use crate::dispatcher::Dispatcher;
121 use crate::environment::EnvironmentDrivers;
122 use crate::graph::{graph_opts, GraphOptions};
123 use crate::node::Status;
124
125 fn run_tokio_local<F>(future: F) -> F::Output
126 where
127 F: std::future::Future,
128 {
129 let runtime = tokio::runtime::Builder::new_current_thread()
130 .enable_all()
131 .build()
132 .expect("tokio current-thread runtime");
133 let local = tokio::task::LocalSet::new();
134 local.block_on(&runtime, future)
135 }
136
137 async fn wait_until(label: &str, mut done: impl FnMut() -> bool) {
138 tokio::time::timeout(Duration::from_secs(5), async {
139 while !done() {
140 tokio::task::yield_now().await;
141 }
142 })
143 .await
144 .unwrap_or_else(|_| panic!("timed out waiting for {label}"));
145 }
146
147 #[test]
148 fn worker_derived_runs_owned_compute_off_graph_thread_and_emits_later() {
149 run_tokio_local(async {
150 let graph_thread = std::thread::current().id();
151 let worker_thread = Arc::new(Mutex::new(None::<ThreadId>));
152 let worker_thread_for_compute = worker_thread.clone();
153 let g = graph_opts(GraphOptions {
154 environment: EnvironmentDrivers::new().with_local_async(Rc::new(TokioLocalDriver)),
155 ..GraphOptions::default()
156 });
157 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
158 let doubled = worker_derived(
159 &g,
160 vec![source.erased()],
161 |ctx| ctx.data::<i32>(0).map(|v| *v),
162 move |value| {
163 *worker_thread_for_compute
164 .lock()
165 .expect("worker thread lock") = Some(std::thread::current().id());
166 Ok::<_, String>(value * 2)
167 },
168 GraphNodeOpts::named("worker"),
169 );
170 let _sub = doubled.subscribe(|_| {});
171
172 source.set(21);
173
174 assert_eq!(doubled.cache(), None);
175 wait_until("worker result", || doubled.cache() == Some(42)).await;
176 assert_ne!(
177 worker_thread
178 .lock()
179 .expect("worker thread lock")
180 .expect("worker ran"),
181 graph_thread
182 );
183
184 let snap = g.describe();
185 assert!(snap
186 .edges
187 .iter()
188 .any(|edge| edge.from == "source" && edge.to == "worker"));
189 assert!(snap.nodes.iter().any(|node| node.id == "worker"
190 && node.factory == "workerDerived"
191 && node.status == Status::Settled));
192 });
193 }
194
195 struct NeverDriver;
196
197 impl LocalAsyncDriver for NeverDriver {
198 fn sleep(&self, _duration: Duration, _callback: Box<dyn FnOnce()>) -> DriverCancel {
199 panic!("NeverDriver.sleep should not be called")
200 }
201
202 fn interval(&self, _period: Duration, _callback: Rc<dyn Fn()>) -> DriverCancel {
203 panic!("NeverDriver.interval should not be called")
204 }
205
206 fn spawn_local(&self, _fut: Pin<Box<dyn Future<Output = ()> + 'static>>) -> DriverCancel {
207 panic!("NeverDriver.spawn_local should not be called")
208 }
209 }
210
211 struct PanickingSpawnDriver;
212
213 impl LocalAsyncDriver for PanickingSpawnDriver {
214 fn sleep(&self, _duration: Duration, _callback: Box<dyn FnOnce()>) -> DriverCancel {
215 panic!("PanickingSpawnDriver.sleep should not be called")
216 }
217
218 fn interval(&self, _period: Duration, _callback: Rc<dyn Fn()>) -> DriverCancel {
219 panic!("PanickingSpawnDriver.interval should not be called")
220 }
221
222 fn spawn_local(&self, _fut: Pin<Box<dyn Future<Output = ()> + 'static>>) -> DriverCancel {
223 panic!("local waiter unavailable")
224 }
225 }
226
227 #[test]
228 fn worker_derived_routes_compute_error_as_later_error_wave() {
229 run_tokio_local(async {
230 let g = graph_opts(GraphOptions {
231 environment: EnvironmentDrivers::new().with_local_async(Rc::new(TokioLocalDriver)),
232 ..GraphOptions::default()
233 });
234 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
235 let worker = worker_derived(
236 &g,
237 vec![source.erased()],
238 |ctx| ctx.data::<i32>(0).map(|v| *v),
239 |_value| Err::<i32, _>("worker failed"),
240 GraphNodeOpts::named("worker"),
241 );
242 let errors = Rc::new(RefCell::new(Vec::new()));
243 let errors_sink = errors.clone();
244 let _sub = worker.subscribe(move |msg| {
245 if let Message::Error(error) = msg {
246 errors_sink.borrow_mut().push(error.to_string());
247 }
248 });
249
250 source.set(1);
251
252 wait_until("worker error", || worker.status() == Status::Errored).await;
253 assert_eq!(&*errors.borrow(), &["worker failed".to_owned()]);
254 });
255 }
256
257 #[test]
258 fn worker_derived_routes_worker_panic_as_later_error_wave() {
259 run_tokio_local(async {
260 let g = graph_opts(GraphOptions {
261 environment: EnvironmentDrivers::new().with_local_async(Rc::new(TokioLocalDriver)),
262 ..GraphOptions::default()
263 });
264 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
265 let worker = worker_derived(
266 &g,
267 vec![source.erased()],
268 |ctx| ctx.data::<i32>(0).map(|v| *v),
269 |_value| -> Result<i32, String> { panic!("worker boom") },
270 GraphNodeOpts::named("worker"),
271 );
272 let errors = Rc::new(RefCell::new(Vec::new()));
273 let errors_sink = errors.clone();
274 let _sub = worker.subscribe(move |msg| {
275 if let Message::Error(error) = msg {
276 errors_sink.borrow_mut().push(error.to_string());
277 }
278 });
279
280 source.set(1);
281
282 wait_until("worker panic error", || worker.status() == Status::Errored).await;
283 assert_eq!(errors.borrow().len(), 1);
284 assert!(errors.borrow()[0].contains("worker_derived: worker task failed"));
285 });
286 }
287
288 #[test]
289 fn worker_derived_drops_superseded_worker_results() {
290 run_tokio_local(async {
291 let g = graph_opts(GraphOptions {
292 environment: EnvironmentDrivers::new().with_local_async(Rc::new(TokioLocalDriver)),
293 ..GraphOptions::default()
294 });
295 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
296 let worker = worker_derived(
297 &g,
298 vec![source.erased()],
299 |ctx| ctx.data::<i32>(0).map(|v| *v),
300 |value| {
301 if value == 1 {
302 std::thread::sleep(Duration::from_millis(75));
303 }
304 Ok::<_, String>(value)
305 },
306 GraphNodeOpts::named("worker"),
307 );
308 let _sub = worker.subscribe(|_| {});
309
310 source.set(1);
311 source.set(2);
312
313 wait_until("latest worker result", || worker.cache() == Some(2)).await;
314 tokio::time::sleep(Duration::from_millis(120)).await;
315 assert_eq!(worker.cache(), Some(2));
316 });
317 }
318
319 #[test]
320 fn worker_derived_no_submit_invocation_fences_prior_worker_result() {
321 run_tokio_local(async {
322 let g = graph_opts(GraphOptions {
323 environment: EnvironmentDrivers::new().with_local_async(Rc::new(TokioLocalDriver)),
324 ..GraphOptions::default()
325 });
326 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
327 let worker = worker_derived(
328 &g,
329 vec![source.erased()],
330 |ctx| {
331 let value = *ctx.data::<i32>(0)?;
332 (value != 0).then_some(value)
333 },
334 |value| {
335 std::thread::sleep(Duration::from_millis(75));
336 Ok::<_, String>(value)
337 },
338 GraphNodeOpts::named("worker"),
339 );
340 let _sub = worker.subscribe(|_| {});
341
342 source.set(1);
343 source.set(0);
344
345 tokio::time::sleep(Duration::from_millis(120)).await;
346 assert_eq!(worker.cache(), None);
347 assert_eq!(worker.status(), Status::Sentinel);
348 });
349 }
350
351 #[test]
352 fn worker_derived_dep_complete_does_not_seal_pending_worker_result() {
353 run_tokio_local(async {
354 let g = graph_opts(GraphOptions {
355 environment: EnvironmentDrivers::new().with_local_async(Rc::new(TokioLocalDriver)),
356 ..GraphOptions::default()
357 });
358 let source = g.producer_opts::<i32, _>(|_ctx| {}, GraphNodeOpts::named("source"));
359 let worker = worker_derived(
360 &g,
361 vec![source.erased()],
362 |ctx| ctx.data::<i32>(0).map(|v| *v),
363 |value| {
364 std::thread::sleep(Duration::from_millis(25));
365 Ok::<_, String>(value * 3)
366 },
367 GraphNodeOpts::named("worker"),
368 );
369 let _sub = worker.subscribe(|_| {});
370
371 source.down(vec![Message::Data(Rc::new(7i32)), Message::Complete]);
372
373 wait_until("worker result after dep complete", || {
374 worker.cache() == Some(21) && worker.status() == Status::Settled
375 })
376 .await;
377 });
378 }
379
380 #[test]
381 fn worker_derived_missing_local_async_driver_errors_on_activation() {
382 let g = graph_opts(GraphOptions::default());
383 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
384 let worker = worker_derived(
385 &g,
386 vec![source.erased()],
387 |ctx| ctx.data::<i32>(0).map(|v| *v),
388 Ok::<_, String>,
389 GraphNodeOpts::named("worker"),
390 );
391 let errors = Rc::new(RefCell::new(Vec::new()));
392 let errors_sink = errors.clone();
393 let _sub = worker.subscribe(move |msg| {
394 if let Message::Error(error) = msg {
395 errors_sink.borrow_mut().push(error.to_string());
396 }
397 });
398
399 source.set(1);
400
401 assert_eq!(
402 &*errors.borrow(),
403 &["worker_derived: missing local async driver".to_owned()]
404 );
405 }
406
407 #[test]
408 fn worker_derived_missing_worker_backend_errors_before_spawning_local_task() {
409 let dispatcher = Dispatcher::new();
410 dispatcher.set_worker_backend_for_test(false);
411 let g = graph_opts(GraphOptions {
412 dispatcher: Some(dispatcher),
413 environment: EnvironmentDrivers::new().with_local_async(Rc::new(NeverDriver)),
414 ..GraphOptions::default()
415 });
416 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
417 let worker = worker_derived(
418 &g,
419 vec![source.erased()],
420 |ctx| ctx.data::<i32>(0).map(|v| *v),
421 Ok::<_, String>,
422 GraphNodeOpts::named("worker"),
423 );
424 let errors = Rc::new(RefCell::new(Vec::new()));
425 let errors_sink = errors.clone();
426 let _sub = worker.subscribe(move |msg| {
427 if let Message::Error(error) = msg {
428 errors_sink.borrow_mut().push(error.to_string());
429 }
430 });
431
432 source.set(1);
433
434 assert_eq!(
435 &*errors.borrow(),
436 &["worker_derived: missing worker backend".to_owned()]
437 );
438 }
439
440 #[test]
441 fn worker_derived_does_not_start_worker_before_local_waiter_is_scheduled() {
442 run_tokio_local(async {
443 let compute_runs = Arc::new(AtomicUsize::new(0));
444 let compute_runs_for_worker = compute_runs.clone();
445 let g = graph_opts(GraphOptions {
446 environment: EnvironmentDrivers::new()
447 .with_local_async(Rc::new(PanickingSpawnDriver)),
448 ..GraphOptions::default()
449 });
450 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
451 let worker = worker_derived(
452 &g,
453 vec![source.erased()],
454 |ctx| ctx.data::<i32>(0).map(|v| *v),
455 move |value| {
456 compute_runs_for_worker.fetch_add(1, Ordering::SeqCst);
457 Ok::<_, String>(value)
458 },
459 GraphNodeOpts::named("worker"),
460 );
461 let _sub = worker.subscribe(|_| {});
462
463 source.set(1);
464 tokio::task::yield_now().await;
465
466 assert_eq!(compute_runs.load(Ordering::SeqCst), 0);
467 });
468 }
469
470 #[test]
471 fn worker_derived_missing_tokio_runtime_errors_before_spawning_local_task() {
472 let g = graph_opts(GraphOptions {
473 environment: EnvironmentDrivers::new().with_local_async(Rc::new(NeverDriver)),
474 ..GraphOptions::default()
475 });
476 let source = g.state_empty_opts::<i32>(GraphNodeOpts::named("source"));
477 let worker = worker_derived(
478 &g,
479 vec![source.erased()],
480 |ctx| ctx.data::<i32>(0).map(|v| *v),
481 Ok::<_, String>,
482 GraphNodeOpts::named("worker"),
483 );
484 let errors = Rc::new(RefCell::new(Vec::new()));
485 let errors_sink = errors.clone();
486 let _sub = worker.subscribe(move |msg| {
487 if let Message::Error(error) = msg {
488 errors_sink.borrow_mut().push(error.to_string());
489 }
490 });
491
492 source.set(1);
493
494 assert_eq!(
495 &*errors.borrow(),
496 &["worker_derived: missing Tokio runtime".to_owned()]
497 );
498 }
499}