Skip to main content

graphrefly/
worker.rs

1//! Graph-helper-first worker compute (D137/D138).
2//!
3//! Worker work never receives `Ctx`, `Node`, `Rc<RefCell<...>>`, live topology,
4//! or erased graph values. The graph-thread kickoff prepares one owned `Send`
5//! input from normal ctx dep reads, then submits the owned compute to the
6//! dispatcher-owned worker backend. Completion is awaited on the graph-local
7//! async driver and emitted through `DeferredCtx` as a fresh later wave.
8
9use 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
21/// Create an async-pool derived node whose CPU-heavy part runs on Tokio's
22/// blocking worker pool.
23///
24/// `prepare` runs synchronously on the graph thread and must return an owned
25/// worker input. `compute` runs off-thread and may only close over `Send + Sync`
26/// state. The result or error comes back as a brand-new graph wave via
27/// `DeferredCtx`, preserving F-SYNC-CORE and D22.
28pub 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            // A dep COMPLETE must not seal the worker before an in-flight owned
49            // compute result returns through DeferredCtx.
50            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}