1use crate::{
48 Result, api,
49 api::system::SystemState,
50 auth::{AuthState, auth_filter, handle_auth_rejection},
51 config::DashboardConfig,
52 security::{AllowedOrigins, same_origin, same_origin_writes},
53 websocket::WebSocketState,
54};
55#[cfg(any(feature = "postgres", feature = "mysql"))]
56use hammerwork::JobQueue;
57use std::{net::SocketAddr, sync::Arc};
58use tokio::sync::RwLock;
59use tracing::{error, info};
60use warp::{Filter, Reply};
61
62pub fn app_routes<Q>(
73 queue: Arc<Q>,
74 auth_state: AuthState,
75 system_state: Arc<RwLock<SystemState>>,
76 websocket_state: Arc<RwLock<WebSocketState>>,
77 static_dir: std::path::PathBuf,
78 allowed_origins: AllowedOrigins,
79) -> impl Filter<Extract = (impl Reply,), Error = std::convert::Infallible> + Clone
80where
81 Q: api::history::JobHistory + 'static,
82{
83 let api_routes = WebDashboard::create_api_routes_static(
84 queue,
85 auth_state.clone(),
86 system_state,
87 allowed_origins.clone(),
88 );
89 let websocket_routes =
90 WebDashboard::create_websocket_routes_static(websocket_state, auth_state, allowed_origins);
91 let static_routes = WebDashboard::create_static_routes_static(static_dir);
92
93 api_routes
94 .or(websocket_routes)
95 .or(static_routes)
96 .recover(handle_auth_rejection)
97}
98
99pub struct WebDashboard {
126 config: DashboardConfig,
127 auth_state: AuthState,
128 websocket_state: Arc<RwLock<WebSocketState>>,
129 allowed_origins: AllowedOrigins,
130 event_manager: Option<Arc<hammerwork::events::EventManager>>,
131}
132
133impl WebDashboard {
134 pub async fn new(config: DashboardConfig) -> Result<Self> {
159 config.validate()?;
160 let allowed_origins = AllowedOrigins::new(&config.allowed_origins)?;
161 let auth_state = AuthState::new(config.auth.clone());
162 let websocket_state = Arc::new(RwLock::new(WebSocketState::new(config.websocket.clone())));
163
164 Ok(Self {
165 config,
166 auth_state,
167 websocket_state,
168 allowed_origins,
169 event_manager: None,
170 })
171 }
172
173 pub fn with_event_manager(
181 mut self,
182 event_manager: Arc<hammerwork::events::EventManager>,
183 ) -> Self {
184 self.event_manager = Some(event_manager);
185 self
186 }
187
188 pub async fn start(self) -> Result<()> {
194 match DatabaseBackend::from_url(&self.config.database_url)? {
195 #[cfg(feature = "postgres")]
196 DatabaseBackend::Postgres => {
197 let pool = sqlx::postgres::PgPoolOptions::new()
198 .max_connections(self.config.pool_size)
199 .connect(&self.config.database_url)
200 .await?;
201 info!(
202 "Connected to PostgreSQL with {} connections",
203 self.config.pool_size
204 );
205 let queue = job_queue(pool, &encryption_settings()?).await?;
206 self.serve(queue, "PostgreSQL").await
207 }
208 #[cfg(feature = "mysql")]
209 DatabaseBackend::MySql => {
210 let pool = sqlx::mysql::MySqlPoolOptions::new()
211 .max_connections(self.config.pool_size)
212 .connect(&self.config.database_url)
213 .await?;
214 info!(
215 "Connected to MySQL with {} connections",
216 self.config.pool_size
217 );
218 let queue = job_queue(pool, &encryption_settings()?).await?;
219 self.serve(queue, "MySQL").await
220 }
221 }
222 }
223
224 async fn serve<Q>(self, queue: Q, database_type: &str) -> Result<()>
226 where
227 Q: api::history::JobHistory + 'static,
228 {
229 let bind_addr: SocketAddr = self.config.bind_addr().parse()?;
230 let queue = Arc::new(queue);
231
232 let system_state = Arc::new(RwLock::new(SystemState::new(
234 self.config.clone(),
235 database_type.to_string(),
236 self.config.pool_size,
237 )));
238
239 let live_interval = self.config.websocket.live_update_interval;
241 if !live_interval.is_zero() {
242 crate::live::LiveUpdates::new(
243 queue.clone(),
244 self.websocket_state.clone(),
245 self.config.websocket.live_update_max_jobs,
246 )
247 .spawn(live_interval);
248 }
249 if let Some(event_manager) = &self.event_manager {
250 let subscription = event_manager
251 .subscribe(hammerwork::events::EventFilter::new())
252 .await?;
253 crate::live::forward_job_events(self.websocket_state.clone(), subscription);
254 }
255
256 let routes = app_routes(
258 queue,
259 self.auth_state.clone(),
260 system_state,
261 self.websocket_state.clone(),
262 self.config.static_dir.clone(),
263 self.allowed_origins.clone(),
264 );
265
266 info!("Starting web server on {}", bind_addr);
267
268 let auth_state_cleanup = self.auth_state.clone();
270 tokio::spawn(async move {
271 let mut interval = tokio::time::interval(std::time::Duration::from_secs(300)); loop {
273 interval.tick().await;
274 auth_state_cleanup.cleanup_expired_attempts().await;
275 }
276 });
277
278 let websocket_state_ping = self.websocket_state.clone();
280 let ping_interval = self.config.websocket.ping_interval;
281 tokio::spawn(async move {
282 let mut interval = tokio::time::interval(ping_interval);
283 loop {
284 interval.tick().await;
285 let state = websocket_state_ping.read().await;
286 state.ping_all_connections().await;
287 }
288 });
289
290 let websocket_state_broadcast = self.websocket_state.clone();
292 WebSocketState::start_broadcast_listener(websocket_state_broadcast).await?;
293
294 if self.config.enable_cors {
298 let cors = warp::cors()
299 .allow_origins(self.allowed_origins.as_slice().iter().map(String::as_str))
300 .allow_headers(vec!["content-type", "authorization"])
301 .allow_methods(vec!["GET", "POST", "PUT", "DELETE", "OPTIONS"]);
302 warp::serve(routes.with(cors)).run(bind_addr).await;
303 } else {
304 warp::serve(routes).run(bind_addr).await;
305 }
306
307 Ok(())
308 }
309
310 fn create_api_routes_static<Q>(
312 queue: Arc<Q>,
313 auth_state: AuthState,
314 system_state: Arc<RwLock<SystemState>>,
315 allowed_origins: AllowedOrigins,
316 ) -> impl Filter<Extract = impl Reply, Error = warp::Rejection> + Clone
317 where
318 Q: api::history::JobHistory + 'static,
319 {
320 let health = warp::path("health")
322 .and(warp::path::end())
323 .and(warp::get())
324 .map(|| {
325 warp::reply::json(&serde_json::json!({
326 "status": "healthy",
327 "timestamp": chrono::Utc::now().to_rfc3339(),
328 "version": env!("CARGO_PKG_VERSION")
329 }))
330 });
331
332 let api_routes = api::queues::routes(queue.clone())
334 .or(api::jobs::routes(queue.clone()))
335 .or(api::stats::routes(queue.clone(), system_state.clone()))
336 .or(api::system::routes(queue.clone(), system_state))
337 .or(api::archive::archive_routes(queue));
338
339 let authenticated_api = warp::path("api")
340 .and(same_origin_writes(allowed_origins))
341 .and(auth_filter(auth_state))
342 .untuple_one()
343 .and(api_routes);
344
345 health.or(authenticated_api)
346 }
347
348 fn create_websocket_routes_static(
352 websocket_state: Arc<RwLock<WebSocketState>>,
353 auth_state: AuthState,
354 allowed_origins: AllowedOrigins,
355 ) -> impl Filter<Extract = impl Reply, Error = warp::Rejection> + Clone {
356 warp::path("ws")
357 .and(warp::path::end())
358 .and(same_origin(allowed_origins))
359 .and(auth_filter(auth_state))
360 .and(warp::ws())
361 .and(warp::any().map(move || websocket_state.clone()))
362 .and_then(
363 |_: (), ws: warp::ws::Ws, websocket_state: Arc<RwLock<WebSocketState>>| async move {
364 let max_message_size = websocket_state.read().await.config().max_message_size;
365 let ws = ws
366 .max_message_size(max_message_size)
367 .max_frame_size(max_message_size);
368 Ok::<_, warp::Rejection>(ws.on_upgrade(move |socket| async move {
369 if let Err(e) =
370 WebSocketState::serve_connection(websocket_state, socket).await
371 {
372 error!("WebSocket error: {}", e);
373 }
374 }))
375 },
376 )
377 }
378
379 fn create_static_routes_static(
381 static_dir: std::path::PathBuf,
382 ) -> impl Filter<Extract = impl Reply, Error = warp::Rejection> + Clone {
383 let static_files = warp::path("static").and(warp::fs::dir(static_dir.clone()));
385
386 let index = warp::path::end().and(warp::fs::file(static_dir.join("index.html")));
388
389 let spa_routes = not_api_path().and(warp::fs::file(static_dir.join("index.html")));
391
392 index.or(static_files).or(spa_routes)
393 }
394}
395
396fn not_api_path() -> impl Filter<Extract = (), Error = warp::Rejection> + Clone {
399 warp::path::full()
400 .and_then(|path: warp::path::FullPath| async move {
401 let path = path.as_str();
402 let reserved = ["/api", "/ws"]
403 .iter()
404 .any(|prefix| path == *prefix || path.starts_with(&format!("{prefix}/")));
405 if reserved {
406 Err(warp::reject::not_found())
407 } else {
408 Ok(())
409 }
410 })
411 .untuple_one()
412}
413
414#[derive(Debug, Clone, Copy, PartialEq, Eq)]
416enum DatabaseBackend {
417 #[cfg(feature = "postgres")]
418 Postgres,
419 #[cfg(feature = "mysql")]
420 MySql,
421}
422
423impl DatabaseBackend {
424 fn from_url(url: &str) -> Result<Self> {
427 if url.starts_with("postgres://") || url.starts_with("postgresql://") {
428 #[cfg(feature = "postgres")]
429 return Ok(Self::Postgres);
430 #[cfg(not(feature = "postgres"))]
431 return Err(anyhow::anyhow!(
432 "PostgreSQL support not enabled. Rebuild with --features postgres"
433 ));
434 }
435 if url.starts_with("mysql://") {
436 #[cfg(feature = "mysql")]
437 return Ok(Self::MySql);
438 #[cfg(not(feature = "mysql"))]
439 return Err(anyhow::anyhow!(
440 "MySQL support not enabled. Rebuild with --features mysql"
441 ));
442 }
443 Err(anyhow::anyhow!(
444 "Unsupported database URL (expected postgres://, postgresql:// or mysql://)"
445 ))
446 }
447}
448
449#[cfg(any(feature = "postgres", feature = "mysql"))]
454fn encryption_settings() -> Result<hammerwork::config::PayloadEncryptionConfig> {
455 let path = std::env::var("HAMMERWORK_ENCRYPTION_CONFIG")
456 .ok()
457 .filter(|path| !path.is_empty());
458 Ok(hammerwork::config::PayloadEncryptionConfig::load(
459 path.as_deref().map(std::path::Path::new),
460 )?)
461}
462
463#[cfg(any(feature = "postgres", feature = "mysql"))]
471async fn job_queue<DB: hammerwork::encryption::KeyManagerBackend>(
472 pool: sqlx::Pool<DB>,
473 encryption: &hammerwork::config::PayloadEncryptionConfig,
474) -> Result<JobQueue<DB>> {
475 let queue = JobQueue::new(pool)
476 .apply_encryption_config(encryption)
477 .await
478 .map_err(|e| {
479 anyhow::anyhow!(
480 "Cannot set up payload encryption from the application's [encryption] \
481 settings ({}); refusing to start rather than store jobs in plaintext",
482 e
483 )
484 })?;
485 if encryption.enabled {
486 info!(
487 "Payload encryption enabled for queues: {:?}",
488 encryption.encrypted_queues
489 );
490 }
491 Ok(queue.with_plaintext_guard(true))
492}
493
494#[cfg(test)]
495mod tests {
496 use super::*;
497 use crate::config::DashboardConfig;
498 use tempfile::tempdir;
499
500 #[cfg(feature = "postgres")]
501 #[tokio::test]
502 async fn test_job_queue_encrypts_like_the_application() {
503 use base64::Engine as _;
504 let pool = || {
505 sqlx::postgres::PgPoolOptions::new()
506 .connect_lazy("postgres://nobody:nothing@127.0.0.1:1/none")
507 .unwrap()
508 };
509
510 let queue = job_queue(pool(), &Default::default()).await.unwrap();
512 assert!(queue.encryption_engine().is_none());
513
514 let var = format!("HW_WEB_TEST_KEY_{}", uuid::Uuid::new_v4().simple());
516 let settings = hammerwork::config::PayloadEncryptionConfig {
517 enabled: true,
518 key_source: hammerwork::config::KeySourceRef::parse(&format!("env://{var}")).unwrap(),
519 key_id: Some("web-key".to_string()),
520 encrypted_queues: vec!["payments".to_string()],
521 ..Default::default()
522 };
523 let err = job_queue(pool(), &settings)
524 .await
525 .err()
526 .expect("the key is missing");
527 assert!(err.to_string().contains("refusing to start"), "{err}");
528 unsafe {
530 std::env::set_var(
531 &var,
532 base64::engine::general_purpose::STANDARD.encode([3u8; 32]),
533 )
534 };
535 let queue = job_queue(pool(), &settings).await.unwrap();
536 assert_eq!(queue.encryption_engine().unwrap().key_id(), "web-key");
537 }
538
539 #[tokio::test]
540 async fn test_dashboard_creation() {
541 let temp_dir = tempdir().unwrap();
542 let mut config = DashboardConfig::new().with_static_dir(temp_dir.path().to_path_buf());
543 config.auth.enabled = cfg!(feature = "auth");
544
545 let dashboard = WebDashboard::new(config).await;
546 assert!(dashboard.is_ok());
547 }
548
549 #[test]
550 fn test_database_backend_from_url() {
551 #[cfg(feature = "postgres")]
552 {
553 assert_eq!(
554 DatabaseBackend::from_url("postgres://localhost/db").unwrap(),
555 DatabaseBackend::Postgres
556 );
557 assert_eq!(
558 DatabaseBackend::from_url("postgresql://localhost/db").unwrap(),
559 DatabaseBackend::Postgres
560 );
561 }
562 #[cfg(feature = "mysql")]
563 assert_eq!(
564 DatabaseBackend::from_url("mysql://localhost/db").unwrap(),
565 DatabaseBackend::MySql
566 );
567 assert!(DatabaseBackend::from_url("sqlite://db").is_err());
568 assert!(DatabaseBackend::from_url("postgresx://db").is_err());
569 }
570
571 #[test]
572 fn test_cors_configuration() {
573 let config = DashboardConfig::new().with_cors(true);
574 assert!(config.enable_cors);
575 }
576
577 #[tokio::test]
578 async fn new_rejects_cors_without_origins() {
579 let mut config = DashboardConfig::new().with_cors(true);
580 config.auth.enabled = false;
581 let err = WebDashboard::new(config)
582 .await
583 .err()
584 .expect("CORS without origins is refused");
585 assert!(err.to_string().contains("allowed_origins"), "{err}");
586 }
587
588 type Router = warp::filters::BoxedFilter<(Box<dyn Reply>,)>;
589
590 fn routes(websocket: crate::config::WebSocketConfig) -> (Router, tempfile::TempDir) {
592 let dir = tempdir().unwrap();
593 std::fs::write(dir.path().join("index.html"), "SPA").unwrap();
594 let config = DashboardConfig {
595 websocket: websocket.clone(),
596 ..DashboardConfig::new()
597 };
598 let system_state = Arc::new(RwLock::new(SystemState::new(
599 config.clone(),
600 "PostgreSQL".into(),
601 1,
602 )));
603 let auth = AuthState::new(crate::config::AuthConfig {
604 enabled: false,
605 ..Default::default()
606 });
607 let routes = app_routes(
608 crate::api::test_support::unreachable_queue(),
609 auth,
610 system_state,
611 Arc::new(RwLock::new(WebSocketState::new(websocket))),
612 dir.path().to_path_buf(),
613 AllowedOrigins::new(["https://ops.example.com"]).unwrap(),
614 )
615 .map(|reply| Box::new(reply) as Box<dyn Reply>)
616 .boxed();
617 (routes, dir)
618 }
619
620 fn post_job() -> warp::test::RequestBuilder {
621 let body = r#"{"queue_name": "q", "payload": {}}"#;
622 warp::test::request()
623 .method("POST")
624 .path("/api/jobs")
625 .header("host", "127.0.0.1:8080")
626 .header("content-length", body.len().to_string())
627 .body(body)
628 }
629
630 #[tokio::test]
633 async fn cross_site_writes_are_refused() {
634 let (routes, _dir) = routes(Default::default());
635 let status = |request: warp::test::RequestBuilder| {
636 let routes = routes.clone();
637 async move { request.reply(&routes).await.status().as_u16() }
638 };
639
640 assert_eq!(status(post_job()).await, 415);
642 assert_eq!(
643 status(post_job().header("content-type", "text/plain")).await,
644 415
645 );
646 for site in ["cross-site", "same-site"] {
648 let request = post_job()
649 .header("content-type", "application/json")
650 .header("origin", "http://127.0.0.1:9999")
651 .header("sec-fetch-site", site);
652 assert_eq!(status(request).await, 403, "{site}");
653 }
654 let request = post_job()
655 .header("content-type", "application/json")
656 .header("origin", "http://evil.example");
657 assert_eq!(status(request).await, 403);
658 let bulk = warp::test::request()
659 .method("POST")
660 .path("/api/jobs/bulk")
661 .header("origin", "null")
662 .json(&serde_json::json!({"job_ids": [], "action": "delete"}));
663 assert_eq!(status(bulk).await, 403);
664 let purge = warp::test::request()
665 .method("DELETE")
666 .path("/api/archive/purge")
667 .header("sec-fetch-site", "cross-site")
668 .json(&serde_json::json!({"older_than": "2020-01-01T00:00:00Z", "dry_run": false}));
669 assert_eq!(status(purge).await, 403);
670
671 let own = post_job()
674 .header("content-type", "application/json")
675 .header("origin", "http://127.0.0.1:8080")
676 .header("sec-fetch-site", "same-origin");
677 assert_eq!(status(own).await, 500);
678 let allowed = post_job()
679 .header("content-type", "application/json; charset=utf-8")
680 .header("origin", "https://ops.example.com")
681 .header("sec-fetch-site", "cross-site");
682 assert_eq!(status(allowed).await, 500);
683 let curl = post_job().header("content-type", "application/json");
684 assert_eq!(status(curl).await, 500);
685
686 let read = warp::test::request()
688 .path("/api/jobs")
689 .header("origin", "http://evil.example")
690 .header("sec-fetch-site", "cross-site");
691 assert_eq!(status(read).await, 500);
692 }
693
694 #[tokio::test]
696 async fn oversized_bodies_are_refused() {
697 let (routes, _dir) = routes(Default::default());
698 let payload = "x".repeat(crate::security::MAX_JSON_BODY_BYTES as usize);
699 let response = warp::test::request()
700 .method("POST")
701 .path("/api/jobs")
702 .json(&serde_json::json!({"queue_name": "q", "payload": payload}))
703 .reply(&routes)
704 .await;
705 assert_eq!(response.status(), 413);
706
707 let ids: Vec<String> = (0..=crate::api::jobs::MAX_BULK_JOB_IDS)
708 .map(|_| uuid::Uuid::new_v4().to_string())
709 .collect();
710 let response = warp::test::request()
711 .method("POST")
712 .path("/api/jobs/bulk")
713 .json(&serde_json::json!({"job_ids": ids, "action": "delete"}))
714 .reply(&routes)
715 .await;
716 assert_eq!(response.status(), 400);
717 let body = String::from_utf8_lossy(response.body()).to_string();
718 assert!(body.contains("Too many job IDs"), "{body}");
719 }
720
721 #[tokio::test]
723 async fn websocket_handshakes_from_other_origins_are_refused() {
724 let (routes, _dir) = routes(Default::default());
725 for (origin, site) in [
728 ("http://evil.example", Some("cross-site")),
729 ("http://evil.example", None),
730 ("null", None),
731 ] {
732 let mut hijack = warp::test::ws().path("/ws").header("origin", origin);
733 if let Some(site) = site {
734 hijack = hijack.header("sec-fetch-site", site);
735 }
736 assert!(
737 hijack.handshake(routes.clone()).await.is_err(),
738 "cross-site WebSocket hijacking from {origin} is refused"
739 );
740 }
741 let own = warp::test::ws()
742 .path("/ws")
743 .header("origin", "http://127.0.0.1:8080")
744 .header("sec-fetch-site", "same-origin")
745 .handshake(routes.clone())
746 .await;
747 assert!(own.is_ok());
748 let allowed = warp::test::ws()
749 .path("/ws")
750 .header("origin", "https://ops.example.com")
751 .handshake(routes)
752 .await;
753 assert!(allowed.is_ok());
754 }
755
756 #[tokio::test]
758 async fn websocket_messages_are_size_limited() {
759 let (routes, _dir) = routes(crate::config::WebSocketConfig {
760 max_message_size: 1024,
761 ..Default::default()
762 });
763 let mut client = warp::test::ws()
764 .path("/ws")
765 .handshake(routes.clone())
766 .await
767 .expect("handshake");
768 client.send_text(r#"{"type": "Ping"}"#).await;
769 let reply = tokio::time::timeout(std::time::Duration::from_secs(2), client.recv())
770 .await
771 .expect("a reply")
772 .expect("open socket");
773 assert!(reply.to_str().unwrap().contains("Pong"));
774
775 client.send_text("x".repeat(4096)).await;
776 let end = tokio::time::timeout(std::time::Duration::from_secs(2), client.recv())
777 .await
778 .expect("the server reacts");
779 assert!(
780 end.is_err() || end.as_ref().is_ok_and(|message| message.is_close()),
781 "the oversized message closed the connection: {end:?}"
782 );
783 }
784}