Skip to main content

ora_worker/
worker.rs

1//! Worker implementation.
2
3use std::{
4    collections::HashMap, iter::once, mem, num::NonZeroUsize, pin::pin, sync::Arc, time::Duration,
5};
6
7use futures::{stream::FuturesUnordered, StreamExt, TryStreamExt};
8use ora_common::task::WorkerSelector;
9use parking_lot::Mutex;
10use thiserror::Error;
11use tokio::{
12    sync::{mpsc, oneshot, Semaphore},
13    task::JoinHandle,
14};
15use tokio_util::sync::CancellationToken;
16use tracing::Instrument;
17use uuid::Uuid;
18
19use crate::{
20    registry::{noop::NoopWorkerRegistry, HeartbeatData, WorkerMetadata, WorkerRegistry},
21    store::{ReadyTask, WorkerStore, WorkerStoreEvent},
22    RawHandler, TaskContext,
23};
24
25/// Options for a [`Worker`].
26#[derive(Debug, Clone)]
27pub struct WorkerOptions {
28    /// The amount of concurrent tasks that can be spawned.
29    pub concurrent_tasks: NonZeroUsize,
30    /// The timeout after which a task is forcibly cancelled
31    /// after receiving a cancellation request.
32    pub cancellation_timeout: Duration,
33}
34
35impl Default for WorkerOptions {
36    fn default() -> Self {
37        // Rather conservative by default.
38        Self {
39            concurrent_tasks: NonZeroUsize::new(4).unwrap(),
40            cancellation_timeout: Duration::from_secs(30),
41        }
42    }
43}
44
45/// A worker where workers can be registered
46/// and are executed whenever tasks are ready.
47pub struct Worker<S, R = NoopWorkerRegistry> {
48    store: S,
49    registry: R,
50    metadata: WorkerMetadata,
51    id: Uuid,
52    handlers: HashMap<WorkerSelector, Arc<dyn RawHandler + Send + Sync>>,
53    semaphore: Arc<Semaphore>,
54    running_tasks: Arc<Mutex<HashMap<Uuid, RunningTask>>>,
55    options: WorkerOptions,
56    shutdown_send: mpsc::Sender<oneshot::Sender<()>>,
57    shutdown_recv: mpsc::Receiver<oneshot::Sender<()>>,
58}
59
60impl<S: std::fmt::Debug> std::fmt::Debug for Worker<S> {
61    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
62        f.debug_struct("Worker")
63            .field("store", &self.store)
64            .field("handlers", &self.handlers.keys().collect::<Vec<_>>())
65            .field("semaphore", &self.semaphore)
66            .field("options", &self.options)
67            .finish_non_exhaustive()
68    }
69}
70
71impl<S> Worker<S>
72where
73    S: WorkerStore + 'static,
74{
75    /// Create a new worker with the default options.
76    pub fn new(store: S) -> Self {
77        Self::new_with_options(store, WorkerOptions::default())
78    }
79
80    /// Create a new worker.
81    pub fn new_with_options(store: S, options: WorkerOptions) -> Self {
82        let (send, recv) = mpsc::channel(1);
83        Self {
84            store,
85            registry: NoopWorkerRegistry,
86            metadata: WorkerMetadata::default(),
87            id: Uuid::new_v4(),
88            handlers: HashMap::new(),
89            semaphore: Arc::new(Semaphore::new(options.concurrent_tasks.get())),
90            running_tasks: Arc::default(),
91            options,
92            shutdown_send: send,
93            shutdown_recv: recv,
94        }
95    }
96}
97
98impl<S, R> Worker<S, R> {
99    /// Register a handler for the worker.
100    ///
101    /// # Panics
102    ///
103    /// Panics if a handler was already registered with a matching [`WorkerSelector`].
104    pub fn register_handler(&mut self, worker: Arc<dyn RawHandler + Send + Sync>) -> &mut Self {
105        let selector = worker.selector();
106
107        assert!(
108            !self.handlers.contains_key(worker.selector()),
109            "a worker is already registered with the given selector: {selector:?}"
110        );
111
112        if let Some(task) = worker.supported_task() {
113            self.metadata.supported_tasks.push(task);
114        }
115
116        self.handlers.insert(worker.selector().clone(), worker);
117        self
118    }
119
120    /// Set the registry for this worker.
121    pub fn with_registry<R2>(self, registry: R2) -> Worker<S, R2> {
122        Worker {
123            store: self.store,
124            registry,
125            metadata: self.metadata,
126            id: self.id,
127            handlers: self.handlers,
128            semaphore: self.semaphore,
129            running_tasks: self.running_tasks,
130            options: self.options,
131            shutdown_send: self.shutdown_send,
132            shutdown_recv: self.shutdown_recv,
133        }
134    }
135
136    /// Set the name of this worker.
137    #[must_use]
138    pub fn with_name(mut self, name: impl Into<String>) -> Self {
139        self.metadata.name = Some(name.into());
140        self
141    }
142
143    /// Set the description of this worker.
144    #[must_use]
145    pub fn with_description(mut self, description: impl Into<String>) -> Self {
146        self.metadata.description = Some(description.into());
147        self
148    }
149
150    /// Set the version of this worker.
151    #[must_use]
152    pub fn with_version(mut self, version: impl Into<String>) -> Self {
153        self.metadata.version = Some(version.into());
154        self
155    }
156
157    /// Get a handle to this worker.
158    pub fn handle(&self) -> WorkerHandle {
159        WorkerHandle {
160            chan: self.shutdown_send.clone(),
161        }
162    }
163
164    /// Get the worker selectors for the handlers
165    /// registered with this worker.
166    pub fn handlers(&self) -> impl Iterator<Item = &WorkerSelector> {
167        self.handlers.keys()
168    }
169
170    /// Return the count of handlers registered.
171    #[must_use]
172    pub fn handler_count(&self) -> usize {
173        self.handlers.len()
174    }
175
176    /// Get the metadata for this worker.
177    #[must_use]
178    pub fn metadata(&self) -> &WorkerMetadata {
179        &self.metadata
180    }
181}
182
183impl<S, R> Worker<S, R>
184where
185    S: WorkerStore + 'static,
186    R: WorkerRegistry,
187{
188    /// Run the worker indefinitely.
189    ///
190    /// # Errors
191    ///
192    /// The function returns on any store error.
193    ///
194    /// # Panics
195    ///
196    /// Only panics due to bugs.
197    #[tracing::instrument(skip_all)]
198    #[allow(clippy::too_many_lines)]
199    pub async fn run(mut self) -> Result<(), Error> {
200        /// The timeout for registry operations,
201        /// these are optional and should not block
202        /// the worker for long periods.
203        const REGISTRY_TIMEOUT: Duration = Duration::from_secs(5);
204
205        macro_rules! wait_shutdown_all {
206            ($confirm:expr) => {
207                let running_tasks = mem::take(&mut *self.running_tasks.lock());
208
209                let mut tasks: FuturesUnordered<_> = running_tasks
210                    .into_iter()
211                    .map(|(task_id, task)| {
212                        let task = task;
213                        let store = self.store.clone();
214                        async move {
215                            tracing::warn!(%task_id, "cancelling task due to shutdown");
216                            if let Err(error) =
217                                store.task_cancelled(task_id).await.map_err(store_error)
218                            {
219                                tracing::error!(?error, "failed to cancel task");
220                            }
221
222                            task.context.cancellation.cancel();
223                            let _ = task.handle.await;
224                        }
225                    })
226                    .collect();
227
228                while tasks.next().await.is_some() {}
229
230                let _ = ($confirm).send(());
231                return Ok(());
232            };
233        }
234
235        let selectors = self.handlers.keys().cloned().collect::<Vec<_>>();
236
237        let (rt_errors_send, mut rt_errors_recv) = mpsc::channel::<Error>(1);
238
239        if let Ok(shutdown_confirm) = self.shutdown_recv.try_recv() {
240            let _ = shutdown_confirm.send(());
241            return Ok(());
242        }
243
244        let mut events = pin!(self.store.events(&selectors).await.map_err(store_error)?);
245
246        if let Ok(shutdown_confirm) = self.shutdown_recv.try_recv() {
247            let _ = shutdown_confirm.send(());
248            return Ok(());
249        }
250
251        self.spawn_tasks(
252            self.store
253                .ready_tasks(&selectors)
254                .await
255                .map_err(store_error)?
256                .into_iter(),
257            rt_errors_send.clone(),
258        )
259        .await?;
260
261        if let Ok(shutdown_confirm) = self.shutdown_recv.try_recv() {
262            wait_shutdown_all!(shutdown_confirm);
263        }
264
265        let mut heartbeat_interval =
266            tokio::time::interval(self.registry.heartbeat_interval().try_into().unwrap());
267        heartbeat_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
268
269        let res = loop {
270            tokio::select! {
271                error = rt_errors_recv.recv() => {
272                    break Err(error.unwrap());
273                }
274                Some(shutdown_confirm) = self.shutdown_recv.recv() => {
275                    wait_shutdown_all!(shutdown_confirm);
276                }
277                event = events.try_next() => {
278                    let event = event.map_err(store_error)?.ok_or(Error::UnexpectedEventStreamEnd)?;
279                    match event {
280                        WorkerStoreEvent::TaskReady(task) => {
281                            self.spawn_tasks(once(task), rt_errors_send.clone()).await?;
282                        },
283                        WorkerStoreEvent::TaskCancelled(task_id) => {
284                            if let Some(task) = self.running_tasks.lock().remove(&task_id) {
285                                task.context.cancellation.cancel();
286                            }
287                        }
288                    }
289                }
290                _ = heartbeat_interval.tick() => {
291                    let heartbeat_fut = tokio::time::timeout(
292                        REGISTRY_TIMEOUT,
293                        async {
294                            let res = self.registry.heartbeat(self.id, &HeartbeatData {}).await?;
295
296                            if res.should_register {
297                                self.registry.register_worker(self.id, &self.metadata).await?;
298                            }
299
300                            Result::<(), R::Error>::Ok(())
301                        },
302                    );
303
304                    if self.registry.enabled() {
305                        match heartbeat_fut.await {
306                            Ok(Ok(())) => {}
307                            Ok(Err(error)) => {
308                                let err = Error::Registry(Box::new(error));
309                                tracing::warn!(?err, "heartbeat failed");
310                            }
311                            Err(_) => {
312                                tracing::warn!("timed out while sending registry heartbeat");
313                            }
314                        }
315                    }
316                }
317            }
318        };
319
320        if self.registry.enabled() {
321            match tokio::time::timeout(REGISTRY_TIMEOUT, self.registry.unregister_worker(self.id))
322                .await
323            {
324                Ok(Ok(())) => {}
325                Ok(Err(error)) => {
326                    let err = Error::Registry(Box::new(error));
327                    tracing::warn!(?err, "unregister failed");
328                }
329                Err(_) => {
330                    tracing::warn!("timed out while sending unregister");
331                }
332            }
333        }
334
335        res
336    }
337
338    #[tracing::instrument(skip_all)]
339    async fn spawn_tasks(
340        &mut self,
341        tasks: impl Iterator<Item = ReadyTask>,
342        rt_errors: mpsc::Sender<Error>,
343    ) -> Result<(), Error> {
344        for task in tasks {
345            let worker = self
346                .handlers
347                .get(&task.definition.worker_selector)
348                .ok_or(Error::HandlerNotFound)?
349                .clone();
350
351            let permit = self.semaphore.clone().acquire_owned().await.unwrap();
352
353            let should_run = self
354                .store
355                .select_task(task.id, self.id)
356                .await
357                .map_err(store_error)?;
358
359            if !should_run {
360                tracing::debug!(task_id = %task.id, "dropping task");
361                continue;
362            }
363
364            let context = TaskContext {
365                task_id: task.id,
366                cancellation: CancellationToken::new(),
367            };
368
369            let cancellation_timeout = self.options.cancellation_timeout;
370            let store = self.store.clone();
371            let running_tasks = self.running_tasks.clone();
372
373            let task_span = tracing::info_span!(
374                "run_task",
375                task_id = %task.id,
376                kind = &*task.definition.worker_selector.kind,
377            );
378
379            let ctx = context.clone();
380            let rt_errors = rt_errors.clone();
381
382            let task_handle = tokio::spawn(async move {
383                let _permit = permit;
384
385                let cancellation = ctx.cancellation.clone();
386                let mut worker_fut = worker.run(ctx, task.definition);
387
388                if let Err(error) = store.task_started(task.id).await {
389                    let _ = rt_errors.send(store_error(error)).await;
390                    return;
391                }
392
393                tokio::select! {
394                    _ = cancellation.cancelled() => {
395                        tokio::select! {
396                            _ = tokio::time::sleep(cancellation_timeout) => {}
397                            res = &mut worker_fut => {
398                                match res {
399                                    Ok(output) => {
400                                        if let Err(error) = store.task_succeeded(task.id, output, worker.output_format()).await {
401                                            let _ = rt_errors.send(store_error(error)).await;
402                                        }
403                                    },
404                                    Err(error) => {
405                                        if let Err(error) = store.task_failed(task.id, format!("{error:?}")).await {
406                                            let _ = rt_errors.send(store_error(error)).await;
407                                        }
408                                    }
409                                }
410                            }
411                        }
412                    }
413                    res = &mut worker_fut => {
414                        match res {
415                            Ok(output) => {
416                                if let Err(error) = store.task_succeeded(task.id, output, worker.output_format()).await {
417                                    let _ = rt_errors.send(store_error(error)).await;
418                                }
419                            },
420                            Err(error) => {
421                                if let Err(error) = store.task_failed(task.id, format!("{error:?}")).await {
422                                    let _ = rt_errors.send(store_error(error)).await;
423                                }
424                            }
425                        }
426                    }
427                }
428
429                running_tasks.lock().remove(&task.id);
430            }.instrument(task_span));
431
432            self.running_tasks.lock().insert(
433                task.id,
434                RunningTask {
435                    context,
436                    handle: task_handle,
437                },
438            );
439        }
440        Ok(())
441    }
442}
443
444/// A handle to a worker that can be used for graceful shutdowns.
445#[derive(Debug, Clone)]
446#[must_use]
447pub struct WorkerHandle {
448    chan: mpsc::Sender<oneshot::Sender<()>>,
449}
450
451impl WorkerHandle {
452    /// Shutdown the worker by cancelling all tasks and waiting for them
453    /// to finish.
454    ///
455    /// If the worker does not exist anymore, this is effectively a no-op.
456    /// If the worker is not yet started, this will wait for the worker to start
457    /// and will shut it down immediately.
458    pub async fn shutdown(&self) {
459        let (send, recv) = oneshot::channel();
460        let _ = self.chan.send(send).await;
461        let _ = recv.await;
462    }
463}
464
465/// A worker error.
466#[derive(Debug, Error)]
467pub enum Error {
468    /// A specific handler was not found, but the
469    /// still received the task. This
470    /// is either a bug in the worker selector
471    /// or the store.
472    #[error("received task but no matching handler was found")]
473    HandlerNotFound,
474    /// The store event stream ended unexpectedly.
475    #[error("unexpected end of event stream")]
476    UnexpectedEventStreamEnd,
477    /// A store error.
478    #[error("store error: {0:?}")]
479    Store(Box<dyn std::error::Error + Send + Sync>),
480    /// A registry error.
481    #[error("registry error: {0:?}")]
482    Registry(Box<dyn std::error::Error + Send + Sync>),
483}
484
485struct RunningTask {
486    context: TaskContext,
487    handle: JoinHandle<()>,
488}
489
490fn store_error<E: std::error::Error + Send + Sync + 'static>(error: E) -> Error {
491    Error::Store(Box::new(error))
492}