1use 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#[derive(Debug, Clone)]
27pub struct WorkerOptions {
28 pub concurrent_tasks: NonZeroUsize,
30 pub cancellation_timeout: Duration,
33}
34
35impl Default for WorkerOptions {
36 fn default() -> Self {
37 Self {
39 concurrent_tasks: NonZeroUsize::new(4).unwrap(),
40 cancellation_timeout: Duration::from_secs(30),
41 }
42 }
43}
44
45pub 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 pub fn new(store: S) -> Self {
77 Self::new_with_options(store, WorkerOptions::default())
78 }
79
80 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 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 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 #[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 #[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 #[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 pub fn handle(&self) -> WorkerHandle {
159 WorkerHandle {
160 chan: self.shutdown_send.clone(),
161 }
162 }
163
164 pub fn handlers(&self) -> impl Iterator<Item = &WorkerSelector> {
167 self.handlers.keys()
168 }
169
170 #[must_use]
172 pub fn handler_count(&self) -> usize {
173 self.handlers.len()
174 }
175
176 #[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 #[tracing::instrument(skip_all)]
198 #[allow(clippy::too_many_lines)]
199 pub async fn run(mut self) -> Result<(), Error> {
200 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#[derive(Debug, Clone)]
446#[must_use]
447pub struct WorkerHandle {
448 chan: mpsc::Sender<oneshot::Sender<()>>,
449}
450
451impl WorkerHandle {
452 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#[derive(Debug, Error)]
467pub enum Error {
468 #[error("received task but no matching handler was found")]
473 HandlerNotFound,
474 #[error("unexpected end of event stream")]
476 UnexpectedEventStreamEnd,
477 #[error("store error: {0:?}")]
479 Store(Box<dyn std::error::Error + Send + Sync>),
480 #[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}