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}