Skip to main content

apalis_board_api/framework/
actix.rs

1use std::{marker::PhantomData, str::FromStr};
2
3use actix_web::{
4    HttpResponse, Responder, Scope,
5    web::{self, Data, Json},
6};
7use apalis_core::backend::{
8    Backend, BackendExt, FetchById, Filter, ListAllTasks, ListQueues, ListTasks, ListWorkers,
9    Metrics, TaskSink, codec::Codec,
10};
11use serde::{Serialize, de::DeserializeOwned};
12use tokio::sync::RwLock;
13
14use crate::{
15    fetch_queues,
16    framework::{ApiBuilder, RegisterRoute},
17    get_all_tasks, get_all_workers, get_task_by_id, get_tasks, get_workers, overview, push_task,
18    stats_by_queue,
19};
20
21#[cfg(feature = "ui")]
22use crate::ui::ServeUI;
23
24/// Handler struct for Actix web routes.
25#[derive(Debug, Clone)]
26pub struct Handler<S, T, Compact> {
27    _phantom: PhantomData<(S, T, Compact)>,
28}
29
30impl<S, T, Compact> Handler<S, T, Compact> {
31    /// Get tasks for a specific queue.
32    pub async fn get_tasks(
33        storage: web::Data<RwLock<S>>,
34        query: web::Query<Filter>,
35    ) -> impl Responder
36    where
37        T: Serialize + DeserializeOwned + 'static,
38        S: ListTasks<T> + Send + 'static + BackendExt,
39        S::Context: Serialize + 'static,
40        S::IdType: Serialize + 'static,
41        <S as Backend>::Error: std::error::Error + 'static,
42        S::Codec: Codec<T, Compact = Compact> + 'static,
43        Compact: 'static,
44    {
45        let storage = storage.into_inner();
46        let filter = query.into_inner();
47
48        match get_tasks::<S, T, Compact>(storage, filter).await {
49            Ok(tasks) => HttpResponse::Ok().json(tasks),
50            Err(e) => HttpResponse::InternalServerError().json(e),
51        }
52    }
53
54    /// Get statistics for a specific queue.
55    pub async fn stats_by_queue(storage: web::Data<RwLock<S>>) -> impl Responder
56    where
57        S::Error: std::error::Error,
58        S: Metrics + BackendExt,
59    {
60        let storage = storage.into_inner();
61
62        match stats_by_queue::<S>(storage).await {
63            Ok(stats) => HttpResponse::Ok().json(stats),
64            Err(e) => HttpResponse::InternalServerError().json(e),
65        }
66    }
67
68    /// Get workers for a specific queue.
69    pub async fn get_workers(storage: web::Data<RwLock<S>>) -> impl Responder
70    where
71        S: ListWorkers + BackendExt,
72        S::Error: std::error::Error,
73    {
74        let storage = storage.into_inner();
75
76        match get_workers::<S>(storage).await {
77            Ok(workers) => HttpResponse::Ok().json(workers),
78            Err(e) => HttpResponse::InternalServerError().json(e),
79        }
80    }
81
82    /// Push a new task to the specified queue.
83    pub async fn push_task(task: Json<T>, storage: Data<RwLock<S>>) -> impl Responder
84    where
85        T: Serialize + DeserializeOwned + 'static,
86        S: TaskSink<T> + Send + BackendExt,
87        S::Error: std::error::Error,
88        S::Codec: Codec<T, Compact = Compact>,
89        <<S as BackendExt>::Codec as Codec<T>>::Error: std::error::Error,
90    {
91        match push_task(task.into_inner(), storage.into_inner()).await {
92            Ok(_) => HttpResponse::Ok().finish(),
93            Err(e) => HttpResponse::InternalServerError().json(e),
94        }
95    }
96
97    /// Get a task by its ID.
98    pub async fn get_task_by_id(
99        task_id: web::Path<String>,
100        storage: web::Data<RwLock<S>>,
101    ) -> impl Responder
102    where
103        T: Serialize + DeserializeOwned + 'static,
104        S: FetchById<T> + 'static,
105        S::Context: Serialize,
106        S::IdType: Serialize,
107        S::Context: Serialize,
108        S::Error: std::error::Error,
109        S::IdType: FromStr,
110        <<S as Backend>::IdType as FromStr>::Err: std::error::Error,
111    {
112        let task_id = task_id.into_inner();
113        let storage = storage.into_inner();
114
115        match get_task_by_id::<S, T>(task_id, storage).await {
116            Ok(Some(task)) => HttpResponse::Ok().json(task),
117            Ok(None) => HttpResponse::NotFound().finish(),
118            Err(e) => HttpResponse::InternalServerError().json(e),
119        }
120    }
121
122    /// Get all tasks across all queues.
123    pub async fn get_all_tasks(
124        storage: web::Data<RwLock<S>>,
125        query: web::Query<Filter>,
126    ) -> impl Responder
127    where
128        S: ListAllTasks + Send,
129        S::Context: Serialize,
130        S::IdType: Serialize,
131        S::Compact: Serialize,
132        <S as Backend>::Error: std::error::Error,
133        <<S as BackendExt>::Codec as Codec<<S as Backend>::Args>>::Error: std::error::Error,
134    {
135        let storage = storage.into_inner();
136        let filter = query.into_inner();
137
138        match get_all_tasks::<S>(storage, filter).await {
139            Ok(tasks) => HttpResponse::Ok().json(tasks),
140            Err(e) => HttpResponse::InternalServerError().json(e),
141        }
142    }
143
144    /// Get all workers across all queues.
145    pub async fn get_all_workers(storage: web::Data<RwLock<S>>) -> impl Responder
146    where
147        S: ListWorkers,
148        S::Error: std::error::Error,
149    {
150        let storage = storage.into_inner();
151
152        match get_all_workers::<S>(storage).await {
153            Ok(workers) => HttpResponse::Ok().json(workers),
154            Err(e) => HttpResponse::InternalServerError().json(e),
155        }
156    }
157
158    /// Fetch all queues.
159    pub async fn fetch_queues(storage: web::Data<RwLock<S>>) -> impl Responder
160    where
161        S::Error: std::error::Error,
162        S: ListQueues,
163    {
164        let storage = storage.into_inner();
165
166        match fetch_queues::<S>(storage).await {
167            Ok(queues) => HttpResponse::Ok().json(queues),
168            Err(e) => HttpResponse::InternalServerError().json(e),
169        }
170    }
171
172    /// Get an overview of statistics.
173    pub async fn overview(storage: web::Data<RwLock<S>>) -> impl Responder
174    where
175        S::Error: std::error::Error,
176        S: Metrics,
177    {
178        let storage = storage.into_inner();
179
180        match overview::<S>(storage).await {
181            Ok(stats) => HttpResponse::Ok().json(stats),
182            Err(e) => HttpResponse::InternalServerError().json(e),
183        }
184    }
185}
186
187impl<B, T, Compact> RegisterRoute<B, T> for ApiBuilder<Scope>
188where
189    B: Metrics + ListWorkers + ListAllTasks + ListQueues + Send + 'static,
190    B::Context: Serialize,
191    B::IdType: Serialize,
192    <B as Backend>::Error: std::error::Error,
193    B::IdType: FromStr,
194    <<B as Backend>::IdType as FromStr>::Err: std::error::Error,
195    Compact: Serialize + 'static,
196    B::Compact: Serialize,
197    <B as Backend>::Error: std::error::Error,
198    <<B as BackendExt>::Codec as Codec<<B as Backend>::Args>>::Error: std::error::Error,
199    T: Serialize + DeserializeOwned + 'static,
200    B: ListTasks<T> + FetchById<T>,
201    B::Codec: Codec<T, Compact = Compact>,
202    <<B as BackendExt>::Codec as Codec<T>>::Error: std::error::Error,
203    B: TaskSink<T>,
204{
205    fn register(mut self, backend: B) -> Self {
206        let queue = backend.get_queue();
207        let backend = web::Data::new(RwLock::new(backend));
208        if self.root {
209            #[allow(unused_mut)]
210            let mut router = self
211                .router
212                .app_data(backend.clone())
213                .route(
214                    "/queues",
215                    web::get().to(Handler::<B, (), Compact>::fetch_queues),
216                )
217                .route(
218                    "/tasks",
219                    web::get().to(Handler::<B, (), Compact>::get_all_tasks),
220                )
221                .route(
222                    "/workers",
223                    web::get().to(Handler::<B, (), Compact>::get_all_workers),
224                )
225                .route(
226                    "/overview",
227                    web::get().to(Handler::<B, (), Compact>::overview),
228                );
229
230            #[cfg(feature = "sse")]
231            {
232                router = router.route("/events", web::get().to(sse::new_client));
233            }
234
235            self.router = router;
236        }
237        let scope = self.router.service(
238            Scope::new(&format!("/queues/{queue}"))
239                .app_data(web::Data::new(queue))
240                .app_data(backend)
241                .route("/tasks", web::get().to(Handler::<B, T, Compact>::get_tasks))
242                .route(
243                    "/stats",
244                    web::get().to(Handler::<B, T, Compact>::stats_by_queue),
245                )
246                .route(
247                    "/workers",
248                    web::get().to(Handler::<B, T, Compact>::get_workers),
249                )
250                .route("/tasks", web::put().to(Handler::<B, T, Compact>::push_task)) // Allow add jobs via api
251                .route(
252                    "/tasks/{id}",
253                    web::get().to(Handler::<B, T, Compact>::get_task_by_id),
254                ),
255        );
256
257        Self {
258            router: scope,
259            root: false,
260        }
261    }
262}
263
264#[cfg(feature = "ui")]
265mod ui {
266    use super::ServeUI;
267    use actix_web::{
268        HttpRequest, HttpResponse, HttpResponseBuilder,
269        dev::HttpServiceFactory,
270        http::{StatusCode, header},
271    };
272    impl ServeUI {
273        fn serve_file(path: &str) -> HttpResponse {
274            let mut file = Self::get_file(path);
275            if file.is_none() {
276                // Try fallback to index.html for unknown routes
277                file = Self::get_file("index.html");
278            }
279
280            match file {
281                Some(f) => {
282                    let path_str = f.path().to_str().unwrap_or("");
283                    let mut builder = HttpResponse::Ok();
284                    let mut builder =
285                        builder.insert_header((header::CONTENT_TYPE, Self::content_type(path_str)));
286
287                    if let Some(cache) = Self::cache_control(path_str) {
288                        builder = builder.insert_header((header::CACHE_CONTROL, cache));
289                    }
290
291                    builder.body(f.contents().to_vec())
292                }
293                None => HttpResponseBuilder::new(StatusCode::NOT_FOUND).finish(),
294            }
295        }
296    }
297    impl HttpServiceFactory for ServeUI {
298        fn register(self, config: &mut actix_web::dev::AppService) {
299            let resource = actix_web::Resource::new("/{tail:.*}").route(actix_web::web::get().to(
300                move |req: HttpRequest| async move {
301                    let path = req.match_info().query("tail");
302
303                    Self::serve_file(path)
304                },
305            ));
306            resource.register(config);
307        }
308    }
309}
310
311/// Expose Server-Sent Events (SSE) functionality.
312#[cfg(feature = "sse")]
313pub mod sse {
314    use std::{sync::Arc, time::Duration};
315
316    use crate::sse::TracingBroadcaster;
317    use actix_web::web::*;
318    use actix_web_lab::sse::Event;
319    use futures::StreamExt;
320    use std::sync::Mutex;
321
322    /// Create a new SSE client connection.
323    pub async fn new_client(
324        broadcaster: Data<Arc<Mutex<TracingBroadcaster>>>,
325    ) -> impl actix_web::Responder {
326        let rx = broadcaster.lock().unwrap().new_client();
327
328        actix_web_lab::sse::Sse::from_stream(
329            rx.filter(|s| futures::future::ready(s.as_ref().is_ok_and(|e| e.span.is_some())))
330                .map(|entry| {
331                    match actix_web_lab::sse::Data::new_json(
332                        entry.map_err(actix_web::error::ErrorInternalServerError)?,
333                    ) {
334                        Ok(data) => Ok(Event::Data(data)),
335                        Err(e) => Err(actix_web::error::ErrorInternalServerError(e)),
336                    }
337                }),
338        )
339        .with_keep_alive(Duration::from_secs(60 * 5))
340    }
341}