Skip to main content

hammerwork_web/
websocket.rs

1//! WebSocket implementation for real-time dashboard updates.
2//!
3//! This module provides WebSocket functionality for real-time communication between
4//! the dashboard frontend and backend. It supports connection management, message
5//! broadcasting, and automatic ping/pong for connection health.
6//!
7//! # Message Types
8//!
9//! The WebSocket API supports several message types for different events:
10//!
11//! ```rust
12//! use hammerwork_web::websocket::{ClientMessage, ServerMessage, AlertSeverity};
13//! use serde_json::json;
14//!
15//! // Client messages (sent from browser to server)
16//! let subscribe_msg = ClientMessage::Subscribe {
17//!     event_types: vec!["queue_updates".to_string(), "job_updates".to_string()],
18//! };
19//!
20//! let ping_msg = ClientMessage::Ping;
21//!
22//! // Server messages (sent from server to browser)
23//! let alert_msg = ServerMessage::SystemAlert {
24//!     message: "High error rate detected".to_string(),
25//!     severity: AlertSeverity::Warning,
26//! };
27//!
28//! let pong_msg = ServerMessage::Pong;
29//! ```
30//!
31//! # Connection Management
32//!
33//! ```rust
34//! use hammerwork_web::websocket::WebSocketState;
35//! use hammerwork_web::config::WebSocketConfig;
36//!
37//! let config = WebSocketConfig::default();
38//! let ws_state = WebSocketState::new(config);
39//!
40//! assert_eq!(ws_state.connection_count(), 0);
41//! ```
42
43use 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
55/// The event types a client can subscribe to.
56pub const EVENT_TYPES: [&str; 4] = [
57    "queue_updates",
58    "job_updates",
59    "system_alerts",
60    "archive_events",
61];
62
63/// What a connection wants to receive. A new connection receives every event until it
64/// sends its first `Subscribe` or `Unsubscribe`.
65#[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        // The first explicit choice replaces the default of "everything". Only known event
85        // types are kept, so a client cannot grow the set without bound.
86        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/// WebSocket connection state manager.
106///
107/// Each connection has an outgoing queue of `message_buffer_size` messages; messages for a
108/// client whose queue is full are dropped rather than buffered without bound.
109#[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    /// The configuration the state was created with.
132    pub fn config(&self) -> &WebSocketConfig {
133        &self.config
134    }
135
136    /// Queue `message` for one connection. A full queue drops the message: the client is not
137    /// reading, and buffering for it would grow without bound. Returns whether the message
138    /// was queued.
139    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    /// Serve one WebSocket connection until it closes.
154    ///
155    /// The shared `state` is only locked for the short moments that register the connection,
156    /// handle a client message and unregister it, never while waiting for the client: a
157    /// connection that held the lock for its whole lifetime would block every other connection,
158    /// the ping task and the broadcast listener.
159    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        // Spawn task to handle outgoing messages to this client
184        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        // Handle incoming messages from this client
197        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        // Clean up connection and subscriptions
224        {
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    /// Handle a message from a client
236    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            // Send pong response
256            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            // Pong received - connection is alive
262            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    /// Handle a client action
273    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                // Answer the client that asked, not everyone.
301                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    /// Broadcast a message to all connected clients
312    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        // A closed channel belongs to a connection that is shutting down; its handler
317        // removes it, so a failed send is not an error here.
318        for (&connection_id, sender) in &self.connections {
319            Self::enqueue(connection_id, sender, ws_message.clone());
320        }
321
322        Ok(())
323    }
324
325    /// Broadcast a message to the clients subscribed to `event_type`
326    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    /// Publish an archive event to all connected clients
348    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        // Send to the broadcast channel
400        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    /// Send ping to all connections to keep them alive
410    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    /// Get current connection count
431    pub fn connection_count(&self) -> usize {
432        self.connections.len()
433    }
434
435    /// Start the broadcast listener task
436    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); // Release the lock before spawning the task
442
443            tokio::spawn(async move {
444                while let Some(broadcast_message) = receiver.recv().await {
445                    // Determine the event type for subscription filtering
446                    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                    // Convert broadcast message to server message
459                    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                    // Actually broadcast the message to subscribed clients
514                    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/// Messages sent from client to server
529#[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/// Messages sent from server to client
538#[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/// Internal broadcast messages
583#[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/// Queue statistics for WebSocket updates
626#[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/// Job update information
640#[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/// Alert severity levels
651#[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    /// The next text message as JSON, or `None` if nothing arrives soon.
757    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        // Closing one connection frees only that connection.
801        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        // Give the server a moment to apply the subscription.
827        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        // Unsubscribing removes a type again; adding one brings it back.
860        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        // Unsubscribing from a type while on the default keeps all the others.
888        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        // A protocol-level ping gets a protocol-level pong.
921        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        // The connection survives and still works.
943        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        // The server drops the extra socket without registering it.
963        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        // With nobody left, pinging is a no-op.
993        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        // The receiver can only be taken once.
1004        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        // Not subscribed to archive events: never delivered.
1127        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    /// M12: subscribing keeps only known event types, so a client cannot grow the set.
1167    #[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    /// M12: each connection's outgoing queue holds `message_buffer_size` messages; a client
1178    /// that does not read loses messages instead of growing server memory.
1179    #[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        // Connections get a queue of the configured size.
1195        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}