1use crate::config::WebSocketConfig;
44use chrono::{DateTime, Utc};
45use futures_util::{SinkExt, StreamExt};
46pub use hammerwork::archive::{ArchivalReason, ArchivalStats};
47use serde::{Deserialize, Serialize};
48use std::collections::HashMap;
49use std::sync::Arc;
50use tokio::sync::mpsc;
51use tracing::{debug, error, info, warn};
52use uuid::Uuid;
53use warp::ws::Message;
54
55pub const EVENT_TYPES: [&str; 4] = [
57 "queue_updates",
58 "job_updates",
59 "system_alerts",
60 "archive_events",
61];
62
63#[derive(Debug)]
66struct Subscription {
67 all: bool,
68 types: std::collections::HashSet<String>,
69}
70
71impl Subscription {
72 fn everything() -> Self {
73 Self {
74 all: true,
75 types: std::collections::HashSet::new(),
76 }
77 }
78
79 fn wants(&self, event_type: &str) -> bool {
80 self.all || self.types.contains(event_type)
81 }
82
83 fn subscribe(&mut self, event_types: Vec<String>) {
84 self.all = false;
87 self.types.extend(
88 event_types
89 .into_iter()
90 .filter(|event_type| EVENT_TYPES.contains(&event_type.as_str())),
91 );
92 }
93
94 fn unsubscribe(&mut self, event_types: &[String]) {
95 if self.all {
96 self.all = false;
97 self.types = EVENT_TYPES.iter().map(|t| t.to_string()).collect();
98 }
99 for event_type in event_types {
100 self.types.remove(event_type);
101 }
102 }
103}
104
105#[derive(Debug)]
110pub struct WebSocketState {
111 config: WebSocketConfig,
112 connections: HashMap<Uuid, mpsc::Sender<Message>>,
113 subscriptions: HashMap<Uuid, Subscription>,
114 broadcast_sender: mpsc::UnboundedSender<BroadcastMessage>,
115 broadcast_receiver: Option<mpsc::UnboundedReceiver<BroadcastMessage>>,
116}
117
118impl WebSocketState {
119 pub fn new(config: WebSocketConfig) -> Self {
120 let (broadcast_sender, broadcast_receiver) = mpsc::unbounded_channel();
121
122 Self {
123 config,
124 connections: HashMap::new(),
125 subscriptions: HashMap::new(),
126 broadcast_sender,
127 broadcast_receiver: Some(broadcast_receiver),
128 }
129 }
130
131 pub fn config(&self) -> &WebSocketConfig {
133 &self.config
134 }
135
136 fn enqueue(connection_id: Uuid, sender: &mpsc::Sender<Message>, message: Message) -> bool {
140 match sender.try_send(message) {
141 Ok(()) => true,
142 Err(mpsc::error::TrySendError::Full(_)) => {
143 debug!(
144 "WebSocket client {} is not keeping up; dropped a message",
145 connection_id
146 );
147 false
148 }
149 Err(mpsc::error::TrySendError::Closed(_)) => false,
150 }
151 }
152
153 pub async fn serve_connection(
160 state: Arc<tokio::sync::RwLock<WebSocketState>>,
161 websocket: warp::ws::WebSocket,
162 ) -> crate::Result<()> {
163 let connection_id = Uuid::new_v4();
164 let (mut ws_sender, mut ws_receiver) = websocket.split();
165 let (tx, mut rx) = {
166 let guard = state.read().await;
167 mpsc::channel::<Message>(guard.config.message_buffer_size.max(1))
168 };
169
170 {
171 let mut guard = state.write().await;
172 if guard.connections.len() >= guard.config.max_connections {
173 warn!("Maximum WebSocket connections reached, rejecting new connection");
174 return Ok(());
175 }
176 guard.connections.insert(connection_id, tx);
177 guard
178 .subscriptions
179 .insert(connection_id, Subscription::everything());
180 }
181 info!("WebSocket connection established: {}", connection_id);
182
183 let writer = tokio::spawn(async move {
185 while let Some(message) = rx.recv().await {
186 if let Err(e) = ws_sender.send(message).await {
187 debug!(
188 "Failed to send WebSocket message to {}: {}",
189 connection_id, e
190 );
191 break;
192 }
193 }
194 });
195
196 while let Some(result) = ws_receiver.next().await {
198 match result {
199 Ok(message) => {
200 let close = message.is_close();
201 let outcome = {
202 let mut guard = state.write().await;
203 guard.handle_client_message(connection_id, message).await
204 };
205 if let Err(e) = outcome {
206 error!(
207 "Error handling client message from {}: {}",
208 connection_id, e
209 );
210 break;
211 }
212 if close {
213 break;
214 }
215 }
216 Err(e) => {
217 debug!("WebSocket error for connection {}: {}", connection_id, e);
218 break;
219 }
220 }
221 }
222
223 {
225 let mut guard = state.write().await;
226 guard.connections.remove(&connection_id);
227 guard.subscriptions.remove(&connection_id);
228 }
229 writer.abort();
230 info!("WebSocket connection closed: {}", connection_id);
231
232 Ok(())
233 }
234
235 async fn handle_client_message(
237 &mut self,
238 connection_id: Uuid,
239 message: Message,
240 ) -> crate::Result<()> {
241 if message.is_text() {
242 if let Ok(text) = message.to_str() {
243 if let Ok(client_message) = serde_json::from_str::<ClientMessage>(text) {
244 debug!(
245 "Received message from {}: {:?}",
246 connection_id, client_message
247 );
248 self.handle_client_action(connection_id, client_message)
249 .await?;
250 } else {
251 warn!("Invalid message format from {}: {}", connection_id, text);
252 }
253 }
254 } else if message.is_ping() {
255 if let Some(sender) = self.connections.get(&connection_id) {
257 let pong_msg = Message::pong(message.as_bytes().to_vec());
258 Self::enqueue(connection_id, sender, pong_msg);
259 }
260 } else if message.is_pong() {
261 debug!("Pong received from {}", connection_id);
263 } else if message.is_close() {
264 debug!("Close message received from {}", connection_id);
265 } else if message.is_binary() {
266 warn!("Binary message not supported from {}", connection_id);
267 }
268
269 Ok(())
270 }
271
272 async fn handle_client_action(
274 &mut self,
275 connection_id: Uuid,
276 message: ClientMessage,
277 ) -> crate::Result<()> {
278 match message {
279 ClientMessage::Subscribe { event_types } => {
280 info!(
281 "Client {} subscribed to events: {:?}",
282 connection_id, event_types
283 );
284 self.subscriptions
285 .entry(connection_id)
286 .or_insert_with(Subscription::everything)
287 .subscribe(event_types);
288 }
289 ClientMessage::Unsubscribe { event_types } => {
290 info!(
291 "Client {} unsubscribed from events: {:?}",
292 connection_id, event_types
293 );
294 self.subscriptions
295 .entry(connection_id)
296 .or_insert_with(Subscription::everything)
297 .unsubscribe(&event_types);
298 }
299 ClientMessage::Ping => {
300 if let Some(sender) = self.connections.get(&connection_id) {
302 let pong = Message::text(serde_json::to_string(&ServerMessage::Pong)?);
303 Self::enqueue(connection_id, sender, pong);
304 }
305 }
306 }
307
308 Ok(())
309 }
310
311 pub async fn broadcast_to_all(&self, message: ServerMessage) -> crate::Result<()> {
313 let json_message = serde_json::to_string(&message)?;
314 let ws_message = Message::text(json_message);
315
316 for (&connection_id, sender) in &self.connections {
319 Self::enqueue(connection_id, sender, ws_message.clone());
320 }
321
322 Ok(())
323 }
324
325 pub async fn broadcast_to_subscribed(
327 &self,
328 message: ServerMessage,
329 event_type: &str,
330 ) -> crate::Result<()> {
331 let json_message = serde_json::to_string(&message)?;
332 let ws_message = Message::text(json_message);
333
334 for (connection_id, sender) in &self.connections {
335 let wanted = self
336 .subscriptions
337 .get(connection_id)
338 .is_none_or(|subscription| subscription.wants(event_type));
339 if wanted {
340 Self::enqueue(*connection_id, sender, ws_message.clone());
341 }
342 }
343
344 Ok(())
345 }
346
347 pub async fn publish_archive_event(
349 &self,
350 event: hammerwork::archive::ArchiveEvent,
351 ) -> crate::Result<()> {
352 let broadcast_message = match event {
353 hammerwork::archive::ArchiveEvent::JobArchived {
354 job_id,
355 queue,
356 reason,
357 } => BroadcastMessage::JobArchived {
358 job_id: job_id.to_string(),
359 queue,
360 reason,
361 },
362 hammerwork::archive::ArchiveEvent::JobRestored {
363 job_id,
364 queue,
365 restored_by,
366 } => BroadcastMessage::JobRestored {
367 job_id: job_id.to_string(),
368 queue,
369 restored_by,
370 },
371 hammerwork::archive::ArchiveEvent::BulkArchiveStarted {
372 operation_id,
373 estimated_jobs,
374 } => BroadcastMessage::BulkArchiveStarted {
375 operation_id,
376 estimated_jobs,
377 },
378 hammerwork::archive::ArchiveEvent::BulkArchiveProgress {
379 operation_id,
380 jobs_processed,
381 total,
382 } => BroadcastMessage::BulkArchiveProgress {
383 operation_id,
384 jobs_processed,
385 total,
386 },
387 hammerwork::archive::ArchiveEvent::BulkArchiveCompleted {
388 operation_id,
389 stats,
390 } => BroadcastMessage::BulkArchiveCompleted {
391 operation_id,
392 stats,
393 },
394 hammerwork::archive::ArchiveEvent::JobsPurged { count, older_than } => {
395 BroadcastMessage::JobsPurged { count, older_than }
396 }
397 };
398
399 if self.broadcast_sender.send(broadcast_message).is_err() {
401 return Err(anyhow::anyhow!(
402 "Failed to send archive event to broadcast channel"
403 ));
404 }
405
406 Ok(())
407 }
408
409 pub async fn ping_all_connections(&self) {
411 let ping_message = Message::ping(b"ping".to_vec());
412 let mut disconnected = Vec::new();
413
414 for (&connection_id, sender) in &self.connections {
415 if sender.is_closed() {
416 disconnected.push(connection_id);
417 } else {
418 Self::enqueue(connection_id, sender, ping_message.clone());
419 }
420 }
421
422 if !disconnected.is_empty() {
423 debug!(
424 "Detected {} disconnected WebSocket clients during ping",
425 disconnected.len()
426 );
427 }
428 }
429
430 pub fn connection_count(&self) -> usize {
432 self.connections.len()
433 }
434
435 pub async fn start_broadcast_listener(
437 state: Arc<tokio::sync::RwLock<WebSocketState>>,
438 ) -> crate::Result<()> {
439 let mut state_guard = state.write().await;
440 if let Some(mut receiver) = state_guard.broadcast_receiver.take() {
441 drop(state_guard); tokio::spawn(async move {
444 while let Some(broadcast_message) = receiver.recv().await {
445 let event_type = match &broadcast_message {
447 BroadcastMessage::QueueUpdate { .. } => "queue_updates",
448 BroadcastMessage::JobUpdate { .. } => "job_updates",
449 BroadcastMessage::SystemAlert { .. } => "system_alerts",
450 BroadcastMessage::JobArchived { .. } => "archive_events",
451 BroadcastMessage::JobRestored { .. } => "archive_events",
452 BroadcastMessage::BulkArchiveStarted { .. } => "archive_events",
453 BroadcastMessage::BulkArchiveProgress { .. } => "archive_events",
454 BroadcastMessage::BulkArchiveCompleted { .. } => "archive_events",
455 BroadcastMessage::JobsPurged { .. } => "archive_events",
456 };
457
458 let server_message = match broadcast_message {
460 BroadcastMessage::QueueUpdate { queue_name, stats } => {
461 ServerMessage::QueueUpdate { queue_name, stats }
462 }
463 BroadcastMessage::JobUpdate { job } => ServerMessage::JobUpdate { job },
464 BroadcastMessage::SystemAlert { message, severity } => {
465 ServerMessage::SystemAlert { message, severity }
466 }
467 BroadcastMessage::JobArchived {
468 job_id,
469 queue,
470 reason,
471 } => ServerMessage::JobArchived {
472 job_id,
473 queue,
474 reason,
475 },
476 BroadcastMessage::JobRestored {
477 job_id,
478 queue,
479 restored_by,
480 } => ServerMessage::JobRestored {
481 job_id,
482 queue,
483 restored_by,
484 },
485 BroadcastMessage::BulkArchiveStarted {
486 operation_id,
487 estimated_jobs,
488 } => ServerMessage::BulkArchiveStarted {
489 operation_id,
490 estimated_jobs,
491 },
492 BroadcastMessage::BulkArchiveProgress {
493 operation_id,
494 jobs_processed,
495 total,
496 } => ServerMessage::BulkArchiveProgress {
497 operation_id,
498 jobs_processed,
499 total,
500 },
501 BroadcastMessage::BulkArchiveCompleted {
502 operation_id,
503 stats,
504 } => ServerMessage::BulkArchiveCompleted {
505 operation_id,
506 stats,
507 },
508 BroadcastMessage::JobsPurged { count, older_than } => {
509 ServerMessage::JobsPurged { count, older_than }
510 }
511 };
512
513 let state_read = state.read().await;
515 if let Err(e) = state_read
516 .broadcast_to_subscribed(server_message, event_type)
517 .await
518 {
519 error!("Failed to broadcast message: {}", e);
520 }
521 }
522 });
523 }
524 Ok(())
525 }
526}
527
528#[derive(Debug, Deserialize)]
530#[serde(tag = "type")]
531pub enum ClientMessage {
532 Subscribe { event_types: Vec<String> },
533 Unsubscribe { event_types: Vec<String> },
534 Ping,
535}
536
537#[derive(Debug, Serialize)]
539#[serde(tag = "type")]
540pub enum ServerMessage {
541 QueueUpdate {
542 queue_name: String,
543 stats: QueueStats,
544 },
545 JobUpdate {
546 job: JobUpdate,
547 },
548 SystemAlert {
549 message: String,
550 severity: AlertSeverity,
551 },
552 JobArchived {
553 job_id: String,
554 queue: String,
555 reason: ArchivalReason,
556 },
557 JobRestored {
558 job_id: String,
559 queue: String,
560 restored_by: Option<String>,
561 },
562 BulkArchiveStarted {
563 operation_id: String,
564 estimated_jobs: u64,
565 },
566 BulkArchiveProgress {
567 operation_id: String,
568 jobs_processed: u64,
569 total: u64,
570 },
571 BulkArchiveCompleted {
572 operation_id: String,
573 stats: ArchivalStats,
574 },
575 JobsPurged {
576 count: u64,
577 older_than: DateTime<Utc>,
578 },
579 Pong,
580}
581
582#[derive(Debug)]
584pub enum BroadcastMessage {
585 QueueUpdate {
586 queue_name: String,
587 stats: QueueStats,
588 },
589 JobUpdate {
590 job: JobUpdate,
591 },
592 SystemAlert {
593 message: String,
594 severity: AlertSeverity,
595 },
596 JobArchived {
597 job_id: String,
598 queue: String,
599 reason: ArchivalReason,
600 },
601 JobRestored {
602 job_id: String,
603 queue: String,
604 restored_by: Option<String>,
605 },
606 BulkArchiveStarted {
607 operation_id: String,
608 estimated_jobs: u64,
609 },
610 BulkArchiveProgress {
611 operation_id: String,
612 jobs_processed: u64,
613 total: u64,
614 },
615 BulkArchiveCompleted {
616 operation_id: String,
617 stats: ArchivalStats,
618 },
619 JobsPurged {
620 count: u64,
621 older_than: DateTime<Utc>,
622 },
623}
624
625#[derive(Debug, Serialize)]
627pub struct QueueStats {
628 pub pending_count: u64,
629 pub running_count: u64,
630 pub completed_count: u64,
631 pub failed_count: u64,
632 pub dead_count: u64,
633 pub throughput_per_minute: f64,
634 pub avg_processing_time_ms: f64,
635 pub error_rate: f64,
636 pub updated_at: chrono::DateTime<chrono::Utc>,
637}
638
639#[derive(Debug, Serialize)]
641pub struct JobUpdate {
642 pub id: String,
643 pub queue_name: String,
644 pub status: String,
645 pub priority: String,
646 pub attempts: i32,
647 pub updated_at: chrono::DateTime<chrono::Utc>,
648}
649
650#[derive(Debug, Serialize)]
652pub enum AlertSeverity {
653 Info,
654 Warning,
655 Error,
656 Critical,
657}
658
659#[cfg(test)]
660mod tests {
661 use super::*;
662 use crate::config::WebSocketConfig;
663
664 #[test]
665 fn test_websocket_state_creation() {
666 let config = WebSocketConfig::default();
667 let state = WebSocketState::new(config);
668 assert_eq!(state.connection_count(), 0);
669 }
670
671 #[test]
672 fn test_client_message_deserialization() {
673 let json = r#"{"type": "Subscribe", "event_types": ["queue_updates", "job_updates"]}"#;
674 let message: ClientMessage = serde_json::from_str(json).unwrap();
675
676 match message {
677 ClientMessage::Subscribe { event_types } => {
678 assert_eq!(event_types.len(), 2);
679 assert!(event_types.contains(&"queue_updates".to_string()));
680 }
681 _ => panic!("Wrong message type"),
682 }
683 }
684
685 #[test]
686 fn test_server_message_serialization() {
687 let message = ServerMessage::SystemAlert {
688 message: "High error rate detected".to_string(),
689 severity: AlertSeverity::Warning,
690 };
691
692 let json = serde_json::to_string(&message).unwrap();
693 assert!(json.contains("type"));
694 assert!(json.contains("SystemAlert"));
695 assert!(json.contains("High error rate detected"));
696 }
697
698 #[tokio::test]
699 async fn test_broadcast_to_all() {
700 let config = WebSocketConfig::default();
701 let state = WebSocketState::new(config);
702
703 let message = ServerMessage::Pong;
704 let result = state.broadcast_to_all(message).await;
705 assert!(result.is_ok());
706 }
707
708 use std::time::Duration;
709 use tokio::sync::RwLock;
710 use warp::Filter;
711
712 type Shared = Arc<RwLock<WebSocketState>>;
713
714 fn ws_route(
715 state: Shared,
716 ) -> impl Filter<Extract = (impl warp::Reply,), Error = warp::Rejection> + Clone {
717 warp::path("ws")
718 .and(warp::ws())
719 .and(warp::any().map(move || state.clone()))
720 .map(|ws: warp::ws::Ws, state: Shared| {
721 ws.on_upgrade(move |socket| async move {
722 let _ = WebSocketState::serve_connection(state, socket).await;
723 })
724 })
725 }
726
727 async fn connect(
728 route: &(
729 impl Filter<Extract = (impl warp::Reply + 'static,), Error = warp::Rejection>
730 + Clone
731 + Send
732 + Sync
733 + 'static
734 ),
735 ) -> warp::test::WsClient {
736 warp::test::ws()
737 .path("/ws")
738 .handshake(route.clone())
739 .await
740 .expect("handshake")
741 }
742
743 async fn wait_for_connections(state: &Shared, expected: usize) {
744 for _ in 0..200 {
745 if state.read().await.connection_count() == expected {
746 return;
747 }
748 tokio::time::sleep(Duration::from_millis(10)).await;
749 }
750 panic!(
751 "expected {expected} connections, have {}",
752 state.read().await.connection_count()
753 );
754 }
755
756 async fn next_json(client: &mut warp::test::WsClient) -> Option<serde_json::Value> {
758 match tokio::time::timeout(Duration::from_millis(400), client.recv()).await {
759 Ok(Ok(message)) if message.is_text() => {
760 Some(serde_json::from_str(message.to_str().unwrap()).unwrap())
761 }
762 Ok(Ok(other)) => panic!("unexpected frame: {other:?}"),
763 Ok(Err(_)) | Err(_) => None,
764 }
765 }
766
767 fn new_state() -> Shared {
768 Arc::new(RwLock::new(WebSocketState::new(WebSocketConfig::default())))
769 }
770
771 fn alert(message: &str) -> ServerMessage {
772 ServerMessage::SystemAlert {
773 message: message.to_string(),
774 severity: AlertSeverity::Warning,
775 }
776 }
777
778 #[tokio::test]
779 async fn several_clients_connect_at_once_and_all_receive_broadcasts() {
780 let state = new_state();
781 let route = ws_route(state.clone());
782 let mut first = connect(&route).await;
783 let mut second = connect(&route).await;
784 let mut third = connect(&route).await;
785 wait_for_connections(&state, 3).await;
786
787 state
788 .read()
789 .await
790 .broadcast_to_all(alert("hello"))
791 .await
792 .unwrap();
793 for client in [&mut first, &mut second, &mut third] {
794 let message = next_json(client).await.expect("broadcast delivered");
795 assert_eq!(message["type"], "SystemAlert");
796 assert_eq!(message["message"], "hello");
797 assert_eq!(message["severity"], "Warning");
798 }
799
800 drop(first);
802 wait_for_connections(&state, 2).await;
803 state
804 .read()
805 .await
806 .broadcast_to_all(alert("again"))
807 .await
808 .unwrap();
809 assert!(next_json(&mut second).await.is_some());
810 assert!(next_json(&mut third).await.is_some());
811 drop((second, third));
812 wait_for_connections(&state, 0).await;
813 }
814
815 #[tokio::test]
816 async fn subscriptions_filter_events_and_new_clients_get_everything() {
817 let state = new_state();
818 let route = ws_route(state.clone());
819 let mut everything = connect(&route).await;
820 let mut archive_only = connect(&route).await;
821 wait_for_connections(&state, 2).await;
822
823 archive_only
824 .send_text(r#"{"type": "Subscribe", "event_types": ["archive_events"]}"#)
825 .await;
826 tokio::time::sleep(Duration::from_millis(100)).await;
828
829 let guard = state.read().await;
830 guard
831 .broadcast_to_subscribed(alert("a system alert"), "system_alerts")
832 .await
833 .unwrap();
834 guard
835 .broadcast_to_subscribed(
836 ServerMessage::JobsPurged {
837 count: 3,
838 older_than: Utc::now(),
839 },
840 "archive_events",
841 )
842 .await
843 .unwrap();
844 drop(guard);
845
846 assert_eq!(
847 next_json(&mut everything).await.unwrap()["type"],
848 "SystemAlert"
849 );
850 assert_eq!(
851 next_json(&mut everything).await.unwrap()["type"],
852 "JobsPurged"
853 );
854 let only = next_json(&mut archive_only).await.unwrap();
855 assert_eq!(only["type"], "JobsPurged", "the alert was filtered out");
856 assert_eq!(only["count"], 3);
857 assert!(next_json(&mut archive_only).await.is_none());
858
859 archive_only
861 .send_text(r#"{"type": "Unsubscribe", "event_types": ["archive_events"]}"#)
862 .await;
863 archive_only
864 .send_text(r#"{"type": "Subscribe", "event_types": ["system_alerts"]}"#)
865 .await;
866 tokio::time::sleep(Duration::from_millis(100)).await;
867 let guard = state.read().await;
868 guard
869 .broadcast_to_subscribed(
870 ServerMessage::JobsPurged {
871 count: 1,
872 older_than: Utc::now(),
873 },
874 "archive_events",
875 )
876 .await
877 .unwrap();
878 guard
879 .broadcast_to_subscribed(alert("now wanted"), "system_alerts")
880 .await
881 .unwrap();
882 drop(guard);
883 let got = next_json(&mut archive_only).await.unwrap();
884 assert_eq!(got["message"], "now wanted");
885 assert!(next_json(&mut archive_only).await.is_none());
886
887 let mut fresh = connect(&route).await;
889 wait_for_connections(&state, 3).await;
890 fresh
891 .send_text(r#"{"type": "Unsubscribe", "event_types": ["job_updates"]}"#)
892 .await;
893 tokio::time::sleep(Duration::from_millis(100)).await;
894 let guard = state.read().await;
895 guard
896 .broadcast_to_subscribed(alert("kept"), "system_alerts")
897 .await
898 .unwrap();
899 guard
900 .broadcast_to_subscribed(alert("dropped"), "job_updates")
901 .await
902 .unwrap();
903 drop(guard);
904 assert_eq!(next_json(&mut fresh).await.unwrap()["message"], "kept");
905 assert!(next_json(&mut fresh).await.is_none());
906 }
907
908 #[tokio::test]
909 async fn a_client_ping_is_answered_to_that_client_only() {
910 let state = new_state();
911 let route = ws_route(state.clone());
912 let mut asker = connect(&route).await;
913 let mut bystander = connect(&route).await;
914 wait_for_connections(&state, 2).await;
915
916 asker.send_text(r#"{"type": "Ping"}"#).await;
917 assert_eq!(next_json(&mut asker).await.unwrap()["type"], "Pong");
918 assert!(next_json(&mut bystander).await.is_none());
919
920 asker.send(warp::ws::Message::ping(b"hi".to_vec())).await;
922 let reply = tokio::time::timeout(Duration::from_secs(1), asker.recv())
923 .await
924 .unwrap()
925 .unwrap();
926 assert!(reply.is_pong());
927 assert_eq!(reply.as_bytes(), b"hi");
928 }
929
930 #[tokio::test]
931 async fn malformed_and_unsupported_messages_are_ignored() {
932 let state = new_state();
933 let route = ws_route(state.clone());
934 let mut client = connect(&route).await;
935 wait_for_connections(&state, 1).await;
936
937 client.send_text("not json at all").await;
938 client.send_text(r#"{"type": "Dance"}"#).await;
939 client.send(warp::ws::Message::binary(vec![1, 2, 3])).await;
940 client.send(warp::ws::Message::pong(b"x".to_vec())).await;
941
942 client.send_text(r#"{"type": "Ping"}"#).await;
944 assert_eq!(next_json(&mut client).await.unwrap()["type"], "Pong");
945 assert_eq!(state.read().await.connection_count(), 1);
946
947 client.send(warp::ws::Message::close()).await;
948 wait_for_connections(&state, 0).await;
949 }
950
951 #[tokio::test]
952 async fn connections_beyond_the_limit_are_turned_away() {
953 let state = Arc::new(RwLock::new(WebSocketState::new(WebSocketConfig {
954 max_connections: 1,
955 ..WebSocketConfig::default()
956 })));
957 let route = ws_route(state.clone());
958 let mut first = connect(&route).await;
959 wait_for_connections(&state, 1).await;
960
961 let mut rejected = connect(&route).await;
962 let end = tokio::time::timeout(Duration::from_secs(1), rejected.recv()).await;
964 assert!(matches!(end, Ok(Err(_))) || matches!(end, Ok(Ok(ref m)) if m.is_close()));
965 assert_eq!(state.read().await.connection_count(), 1);
966
967 state
968 .read()
969 .await
970 .broadcast_to_all(alert("x"))
971 .await
972 .unwrap();
973 assert!(next_json(&mut first).await.is_some());
974 }
975
976 #[tokio::test]
977 async fn the_ping_task_reaches_every_connection() {
978 let state = new_state();
979 let route = ws_route(state.clone());
980 let mut client = connect(&route).await;
981 wait_for_connections(&state, 1).await;
982
983 state.read().await.ping_all_connections().await;
984 let frame = tokio::time::timeout(Duration::from_secs(1), client.recv())
985 .await
986 .unwrap()
987 .unwrap();
988 assert!(frame.is_ping());
989 assert_eq!(frame.as_bytes(), b"ping");
990 drop(client);
991 wait_for_connections(&state, 0).await;
992 state.read().await.ping_all_connections().await;
994 }
995
996 #[tokio::test]
997 async fn archive_events_flow_through_the_broadcast_listener() {
998 use hammerwork::archive::ArchiveEvent;
999 let state = new_state();
1000 WebSocketState::start_broadcast_listener(state.clone())
1001 .await
1002 .unwrap();
1003 WebSocketState::start_broadcast_listener(state.clone())
1005 .await
1006 .unwrap();
1007 let route = ws_route(state.clone());
1008 let mut client = connect(&route).await;
1009 wait_for_connections(&state, 1).await;
1010
1011 let job_id = Uuid::new_v4();
1012 let older_than = Utc::now();
1013 let events = vec![
1014 ArchiveEvent::JobArchived {
1015 job_id,
1016 queue: "q".into(),
1017 reason: ArchivalReason::Manual,
1018 },
1019 ArchiveEvent::JobRestored {
1020 job_id,
1021 queue: "q".into(),
1022 restored_by: Some("me".into()),
1023 },
1024 ArchiveEvent::BulkArchiveStarted {
1025 operation_id: "op".into(),
1026 estimated_jobs: 10,
1027 },
1028 ArchiveEvent::BulkArchiveProgress {
1029 operation_id: "op".into(),
1030 jobs_processed: 5,
1031 total: 10,
1032 },
1033 ArchiveEvent::BulkArchiveCompleted {
1034 operation_id: "op".into(),
1035 stats: ArchivalStats::default(),
1036 },
1037 ArchiveEvent::JobsPurged {
1038 count: 2,
1039 older_than,
1040 },
1041 ];
1042 for event in events {
1043 state
1044 .read()
1045 .await
1046 .publish_archive_event(event)
1047 .await
1048 .unwrap();
1049 }
1050
1051 let mut seen = Vec::new();
1052 for _ in 0..6 {
1053 seen.push(next_json(&mut client).await.expect("event delivered"));
1054 }
1055 let types: Vec<&str> = seen.iter().map(|m| m["type"].as_str().unwrap()).collect();
1056 assert_eq!(
1057 types,
1058 vec![
1059 "JobArchived",
1060 "JobRestored",
1061 "BulkArchiveStarted",
1062 "BulkArchiveProgress",
1063 "BulkArchiveCompleted",
1064 "JobsPurged"
1065 ]
1066 );
1067 assert_eq!(seen[0]["job_id"], job_id.to_string());
1068 assert_eq!(seen[0]["reason"], "Manual");
1069 assert_eq!(seen[1]["restored_by"], "me");
1070 assert_eq!(seen[2]["estimated_jobs"], 10);
1071 assert_eq!(seen[3]["jobs_processed"], 5);
1072 assert_eq!(seen[4]["stats"]["jobs_archived"], 0);
1073 assert_eq!(seen[5]["count"], 2);
1074 }
1075
1076 #[tokio::test]
1077 async fn the_other_broadcast_kinds_are_converted_and_filtered_by_type() {
1078 let state = new_state();
1079 WebSocketState::start_broadcast_listener(state.clone())
1080 .await
1081 .unwrap();
1082 let route = ws_route(state.clone());
1083 let mut queue_only = connect(&route).await;
1084 wait_for_connections(&state, 1).await;
1085 queue_only
1086 .send_text(r#"{"type": "Subscribe", "event_types": ["queue_updates", "job_updates", "system_alerts"]}"#)
1087 .await;
1088 tokio::time::sleep(Duration::from_millis(100)).await;
1089
1090 let sender = state.read().await.broadcast_sender.clone();
1091 let now = Utc::now();
1092 sender
1093 .send(BroadcastMessage::QueueUpdate {
1094 queue_name: "emails".into(),
1095 stats: QueueStats {
1096 pending_count: 1,
1097 running_count: 2,
1098 completed_count: 3,
1099 failed_count: 4,
1100 dead_count: 5,
1101 throughput_per_minute: 6.0,
1102 avg_processing_time_ms: 7.0,
1103 error_rate: 0.5,
1104 updated_at: now,
1105 },
1106 })
1107 .unwrap();
1108 sender
1109 .send(BroadcastMessage::JobUpdate {
1110 job: JobUpdate {
1111 id: "j1".into(),
1112 queue_name: "emails".into(),
1113 status: "Running".into(),
1114 priority: "High".into(),
1115 attempts: 1,
1116 updated_at: now,
1117 },
1118 })
1119 .unwrap();
1120 sender
1121 .send(BroadcastMessage::SystemAlert {
1122 message: "disk".into(),
1123 severity: AlertSeverity::Critical,
1124 })
1125 .unwrap();
1126 sender
1128 .send(BroadcastMessage::JobsPurged {
1129 count: 9,
1130 older_than: now,
1131 })
1132 .unwrap();
1133
1134 let queue = next_json(&mut queue_only).await.unwrap();
1135 assert_eq!(queue["type"], "QueueUpdate");
1136 assert_eq!(queue["queue_name"], "emails");
1137 assert_eq!(queue["stats"]["dead_count"], 5);
1138 let job = next_json(&mut queue_only).await.unwrap();
1139 assert_eq!(
1140 (job["type"].as_str(), job["job"]["id"].as_str()),
1141 (Some("JobUpdate"), Some("j1"))
1142 );
1143 let alert = next_json(&mut queue_only).await.unwrap();
1144 assert_eq!(alert["severity"], "Critical");
1145 assert!(next_json(&mut queue_only).await.is_none());
1146 }
1147
1148 #[test]
1149 fn subscription_rules() {
1150 let mut sub = Subscription::everything();
1151 assert!(sub.wants("anything"));
1152 sub.unsubscribe(&["job_updates".to_string()]);
1153 assert!(!sub.wants("job_updates") && sub.wants("queue_updates"));
1154 assert!(
1155 !sub.wants("custom"),
1156 "after the first change only known types remain"
1157 );
1158
1159 let mut sub = Subscription::everything();
1160 sub.subscribe(vec!["archive_events".to_string()]);
1161 assert!(sub.wants("archive_events") && !sub.wants("queue_updates"));
1162 sub.subscribe(vec!["queue_updates".to_string()]);
1163 assert!(sub.wants("queue_updates") && sub.wants("archive_events"));
1164 }
1165
1166 #[test]
1168 fn subscriptions_ignore_unknown_event_types() {
1169 let mut sub = Subscription::everything();
1170 sub.subscribe((0..10_000).map(|n| format!("junk-{n}")).collect());
1171 sub.subscribe(vec!["job_updates".to_string(), "job_updates".to_string()]);
1172 assert_eq!(sub.types.len(), 1);
1173 assert!(sub.wants("job_updates"));
1174 assert!(!sub.wants("junk-1"));
1175 }
1176
1177 #[tokio::test]
1180 async fn outgoing_messages_are_bounded_per_connection() {
1181 let (sender, mut receiver) = mpsc::channel(2);
1182 let id = Uuid::new_v4();
1183 assert!(WebSocketState::enqueue(id, &sender, Message::text("1")));
1184 assert!(WebSocketState::enqueue(id, &sender, Message::text("2")));
1185 assert!(
1186 !WebSocketState::enqueue(id, &sender, Message::text("3")),
1187 "a full queue drops the message"
1188 );
1189 assert_eq!(receiver.recv().await.unwrap().to_str().unwrap(), "1");
1190 assert!(WebSocketState::enqueue(id, &sender, Message::text("4")));
1191 drop(receiver);
1192 assert!(!WebSocketState::enqueue(id, &sender, Message::text("5")));
1193
1194 let state = Arc::new(RwLock::new(WebSocketState::new(WebSocketConfig {
1196 message_buffer_size: 3,
1197 ..WebSocketConfig::default()
1198 })));
1199 let route = ws_route(state.clone());
1200 let _client = connect(&route).await;
1201 wait_for_connections(&state, 1).await;
1202 let guard = state.read().await;
1203 let queue = guard.connections.values().next().unwrap();
1204 assert_eq!(queue.max_capacity(), 3);
1205 assert_eq!(guard.config().message_buffer_size, 3);
1206 }
1207
1208 #[test]
1209 fn every_server_message_is_tagged_with_its_type() {
1210 let now = Utc::now();
1211 let messages = vec![
1212 (ServerMessage::Pong, "Pong"),
1213 (alert("x"), "SystemAlert"),
1214 (
1215 ServerMessage::JobArchived {
1216 job_id: "j".into(),
1217 queue: "q".into(),
1218 reason: ArchivalReason::Automatic,
1219 },
1220 "JobArchived",
1221 ),
1222 (
1223 ServerMessage::JobRestored {
1224 job_id: "j".into(),
1225 queue: "q".into(),
1226 restored_by: None,
1227 },
1228 "JobRestored",
1229 ),
1230 (
1231 ServerMessage::BulkArchiveStarted {
1232 operation_id: "o".into(),
1233 estimated_jobs: 1,
1234 },
1235 "BulkArchiveStarted",
1236 ),
1237 (
1238 ServerMessage::JobsPurged {
1239 count: 1,
1240 older_than: now,
1241 },
1242 "JobsPurged",
1243 ),
1244 ];
1245 for (message, expected) in messages {
1246 let json: serde_json::Value =
1247 serde_json::from_str(&serde_json::to_string(&message).unwrap()).unwrap();
1248 assert_eq!(json["type"], expected);
1249 }
1250 for severity in [
1251 AlertSeverity::Info,
1252 AlertSeverity::Warning,
1253 AlertSeverity::Error,
1254 AlertSeverity::Critical,
1255 ] {
1256 assert!(serde_json::to_string(&severity).unwrap().starts_with('"'));
1257 }
1258 for text in [
1259 r#"{"type": "Ping"}"#,
1260 r#"{"type": "Unsubscribe", "event_types": []}"#,
1261 ] {
1262 assert!(serde_json::from_str::<ClientMessage>(text).is_ok());
1263 }
1264 assert!(serde_json::from_str::<ClientMessage>(r#"{"type": "Subscribe"}"#).is_err());
1265 }
1266}