Skip to main content

fraiseql_server/realtime/
connections.rs

1//! Connection manager for tracking active realtime `WebSocket` connections.
2//!
3//! Uses `DashMap` for lock-free concurrent access to connection state.
4//! Each connection has an associated event sender channel for pushing
5//! change events from the delivery pipeline.
6
7use std::sync::atomic::{AtomicUsize, Ordering};
8
9use dashmap::DashMap;
10use tokio::sync::{mpsc, oneshot};
11
12/// Unique identifier for a `WebSocket` connection.
13pub type ConnectionId = String;
14
15/// Signal to gracefully close a connection with a specific `WebSocket` close code.
16#[derive(Debug, Clone)]
17pub struct CloseSignal {
18    /// `WebSocket` close code (e.g., 4002 for "slow consumer").
19    pub code:   u16,
20    /// Human-readable close reason.
21    pub reason: String,
22}
23
24/// State for a single active `WebSocket` connection.
25#[derive(Debug, Clone)]
26pub struct ConnectionState {
27    /// Unique connection identifier (UUID v4).
28    pub connection_id: ConnectionId,
29    /// User identifier (from JWT `sub` claim).
30    pub user_id:       String,
31    /// Security context hash for grouping connections with identical RLS context.
32    pub context_hash:  u64,
33    /// Token expiration (Unix timestamp in seconds).
34    pub expires_at:    i64,
35}
36
37impl ConnectionState {
38    /// Create a new connection state.
39    #[must_use]
40    pub const fn new(
41        connection_id: ConnectionId,
42        user_id: String,
43        context_hash: u64,
44        expires_at: i64,
45    ) -> Self {
46        Self {
47            connection_id,
48            user_id,
49            context_hash,
50            expires_at,
51        }
52    }
53}
54
55/// Thread-safe manager for active `WebSocket` connections.
56///
57/// Uses `DashMap` for lock-free concurrent reads and writes.
58/// Each connection has an event sender for delivering change events and a
59/// oneshot close-signal channel for graceful server-initiated disconnection
60/// (e.g., slow consumer policy).
61pub struct ConnectionManager {
62    /// Active connections indexed by connection ID.
63    connections:               DashMap<ConnectionId, ConnectionState>,
64    /// Per-connection event senders (bounded channel for change events).
65    event_senders:             DashMap<ConnectionId, mpsc::Sender<String>>,
66    /// Per-connection oneshot senders for server-initiated close signals.
67    close_senders:             DashMap<ConnectionId, oneshot::Sender<CloseSignal>>,
68    /// Per-connection consecutive drop counters for slow-consumer detection.
69    drop_counts:               DashMap<ConnectionId, AtomicUsize>,
70    /// Maximum consecutive delivery failures before a connection is kicked.
71    max_consecutive_drops:     usize,
72    /// Capacity of per-connection event channels.
73    connection_event_capacity: usize,
74}
75
76impl ConnectionManager {
77    /// Create a new empty connection manager.
78    #[must_use]
79    pub fn new(max_consecutive_drops: usize, connection_event_capacity: usize) -> Self {
80        Self {
81            connections: DashMap::new(),
82            event_senders: DashMap::new(),
83            close_senders: DashMap::new(),
84            drop_counts: DashMap::new(),
85            max_consecutive_drops,
86            connection_event_capacity,
87        }
88    }
89
90    /// Register a new connection and return a receiver for events and a close-signal receiver.
91    ///
92    /// The event receiver should be polled by the connection handler to forward
93    /// change events to the `WebSocket`. The close-signal receiver fires when the
94    /// delivery pipeline detects a slow consumer.
95    #[must_use]
96    pub fn insert(
97        &self,
98        state: ConnectionState,
99    ) -> (mpsc::Receiver<String>, oneshot::Receiver<CloseSignal>) {
100        let (event_tx, event_rx) = mpsc::channel(self.connection_event_capacity);
101        let (close_tx, close_rx) = oneshot::channel();
102        self.event_senders.insert(state.connection_id.clone(), event_tx);
103        self.close_senders.insert(state.connection_id.clone(), close_tx);
104        self.drop_counts.insert(state.connection_id.clone(), AtomicUsize::new(0));
105        self.connections.insert(state.connection_id.clone(), state);
106        (event_rx, close_rx)
107    }
108
109    /// Remove a connection by ID, cleaning up all associated state.
110    pub fn remove(&self, connection_id: &str) {
111        self.connections.remove(connection_id);
112        self.event_senders.remove(connection_id);
113        self.close_senders.remove(connection_id);
114        self.drop_counts.remove(connection_id);
115    }
116
117    /// Total number of active connections.
118    #[must_use]
119    pub fn count(&self) -> usize {
120        self.connections.len()
121    }
122
123    /// Number of connections for a specific security context hash.
124    #[must_use]
125    pub fn count_by_context(&self, context_hash: u64) -> usize {
126        self.connections
127            .iter()
128            .filter(|entry| entry.value().context_hash == context_hash)
129            .count()
130    }
131
132    /// Send a serialized event to a connection's event channel.
133    ///
134    /// On success, resets the per-connection drop counter.
135    /// On failure (channel full), increments the drop counter. If the counter
136    /// reaches `max_consecutive_drops`, a close signal with code 4002 ("slow
137    /// consumer") is sent to the connection handler.
138    ///
139    /// Returns `true` if the event was sent, `false` otherwise.
140    #[must_use]
141    pub fn send_event(&self, connection_id: &str, json: String) -> bool {
142        let sent = self
143            .event_senders
144            .get(connection_id)
145            .is_some_and(|sender| sender.try_send(json).is_ok());
146
147        if sent {
148            if let Some(counter) = self.drop_counts.get(connection_id) {
149                counter.store(0, Ordering::Relaxed);
150            }
151        } else if let Some(counter) = self.drop_counts.get(connection_id) {
152            let new_count = counter.fetch_add(1, Ordering::Relaxed) + 1;
153            if new_count >= self.max_consecutive_drops {
154                // Slow consumer: signal the connection handler to close with 4002.
155                if let Some((_, close_tx)) = self.close_senders.remove(connection_id) {
156                    let _ = close_tx.send(CloseSignal {
157                        code:   4002,
158                        reason: "slow consumer".to_owned(),
159                    });
160                }
161            }
162        }
163
164        sent
165    }
166
167    /// Return the current consecutive drop count for a connection (for testing).
168    #[must_use]
169    pub fn drop_count(&self, connection_id: &str) -> usize {
170        self.drop_counts.get(connection_id).map_or(0, |c| c.load(Ordering::Relaxed))
171    }
172}
173
174impl Default for ConnectionManager {
175    fn default() -> Self {
176        Self::new(50, 256)
177    }
178}