Skip to main content

fraiseql_server/realtime/
server.rs

1//! Realtime `WebSocket` server and configuration.
2//!
3//! `RealtimeServer` manages `WebSocket` connections, authenticates clients,
4//! handles heartbeats and idle timeouts, enforces connection limits, and
5//! processes subscription requests for entity change events.
6
7use std::{collections::HashSet, sync::Arc, time::Duration};
8
9use axum::{
10    extract::{
11        Query, State,
12        ws::{Message, WebSocket, WebSocketUpgrade},
13    },
14    http::StatusCode,
15    response::IntoResponse,
16};
17use futures::{SinkExt, StreamExt, future::BoxFuture};
18use serde::Deserialize;
19use tracing::{debug, info, warn};
20
21use super::{
22    connections::{ConnectionManager, ConnectionState},
23    protocol::{ClientMessage, ServerMessage},
24    subscriptions::{EventKind, SubscriptionDetails, SubscriptionManager, parse_filter},
25};
26
27/// Configuration for the realtime `WebSocket` server.
28#[derive(Debug, Clone)]
29pub struct RealtimeConfig {
30    /// Maximum concurrent connections per security context hash (default: 10).
31    pub max_connections_per_context:  usize,
32    /// Interval between server heartbeat pings (default: 30s).
33    pub heartbeat_interval:           Duration,
34    /// Disconnect after this duration of inactivity (default: 60s).
35    pub idle_timeout:                 Duration,
36    /// Maximum subscriptions per entity across all connections (default: 10,000).
37    pub max_subscriptions_per_entity: usize,
38    /// Bounded channel capacity for the observer-to-pipeline event channel (default: 10,000).
39    pub event_channel_capacity:       usize,
40    /// How often to re-validate JWT tokens (default: same as `heartbeat_interval`).
41    /// Per D14: check JWT on each heartbeat.
42    pub token_revalidation_interval:  Duration,
43    /// Consecutive per-connection delivery failures before kicking the client
44    /// with close code 4002 "slow consumer" (default: 50).
45    pub max_consecutive_drops:        usize,
46    /// Capacity of the per-connection event channel (default: 256).
47    pub connection_event_capacity:    usize,
48}
49
50impl Default for RealtimeConfig {
51    fn default() -> Self {
52        Self {
53            max_connections_per_context:  10,
54            heartbeat_interval:           Duration::from_secs(30),
55            idle_timeout:                 Duration::from_secs(60),
56            max_subscriptions_per_entity: 10_000,
57            event_channel_capacity:       10_000,
58            token_revalidation_interval:  Duration::from_secs(30),
59            max_consecutive_drops:        50,
60            connection_event_capacity:    256,
61        }
62    }
63}
64
65/// Token validator trait for authenticating `WebSocket` connections.
66///
67/// Implementations validate bearer tokens (JWT or otherwise) and return
68/// the validated token information including user identity and expiration.
69///
70/// The trait is object-safe (returns `BoxFuture`) so it can be stored as
71/// `Arc<dyn TokenValidator>` and mounted in `build_base_router` without
72/// adding a type parameter to `Server<A>`.
73///
74/// # Production implementation
75///
76/// `JwtTokenValidator` wraps the existing `OidcValidator`
77/// from `fraiseql-core`, reusing the same JWT validation path as the GraphQL
78/// endpoint's OIDC middleware.
79pub trait TokenValidator: Send + Sync + 'static {
80    /// Validate a bearer token and return the token info.
81    ///
82    /// # Errors
83    ///
84    /// Returns an error string if the token is invalid, expired, or
85    /// cannot be validated.
86    fn validate<'a>(&'a self, token: &'a str) -> BoxFuture<'a, Result<TokenInfo, String>>;
87}
88
89/// Information extracted from a validated token.
90#[derive(Debug, Clone)]
91pub struct TokenInfo {
92    /// User identifier (from JWT `sub` claim).
93    pub user_id:      String,
94    /// Security context hash for connection grouping.
95    pub context_hash: u64,
96    /// When the token expires (Unix timestamp in seconds).
97    pub expires_at:   i64,
98}
99
100/// Shared state for the realtime `WebSocket` handler.
101#[derive(Clone)]
102pub struct RealtimeState {
103    /// The realtime server instance.
104    pub server:    Arc<RealtimeServer>,
105    /// Token validator for authenticating connections.
106    pub validator: Arc<dyn TokenValidator>,
107}
108
109/// The realtime `WebSocket` server.
110pub struct RealtimeServer {
111    /// Active connection manager.
112    pub(crate) connections:    Arc<ConnectionManager>,
113    /// Subscription manager for entity change subscriptions.
114    pub(crate) subscriptions:  Arc<SubscriptionManager>,
115    /// Set of entity names that accept realtime subscriptions.
116    pub(crate) known_entities: HashSet<String>,
117    /// Server configuration.
118    pub(crate) config:         RealtimeConfig,
119}
120
121impl RealtimeServer {
122    /// Create a new realtime server with the given configuration.
123    #[must_use]
124    pub fn new(config: RealtimeConfig) -> Self {
125        let max_subs = config.max_subscriptions_per_entity;
126        let connections = Arc::new(ConnectionManager::new(
127            config.max_consecutive_drops,
128            config.connection_event_capacity,
129        ));
130        Self {
131            connections,
132            subscriptions: Arc::new(SubscriptionManager::new(max_subs)),
133            known_entities: HashSet::new(),
134            config,
135        }
136    }
137
138    /// Create a new realtime server with known entities for subscription validation.
139    #[must_use]
140    pub fn with_entities(config: RealtimeConfig, entities: HashSet<String>) -> Self {
141        let max_subs = config.max_subscriptions_per_entity;
142        let connections = Arc::new(ConnectionManager::new(
143            config.max_consecutive_drops,
144            config.connection_event_capacity,
145        ));
146        Self {
147            connections,
148            subscriptions: Arc::new(SubscriptionManager::new(max_subs)),
149            known_entities: entities,
150            config,
151        }
152    }
153
154    /// Returns the number of active connections.
155    #[must_use]
156    pub fn active_connections(&self) -> usize {
157        self.connections.count()
158    }
159}
160
161/// Query parameters for the `WebSocket` upgrade request.
162#[derive(Debug, Deserialize)]
163pub struct WsQueryParams {
164    /// Bearer token passed as query parameter.
165    pub token: Option<String>,
166}
167
168/// `WebSocket` upgrade handler for `/realtime/v1`.
169///
170/// Authenticates via `?token=` query parameter or `Authorization: Bearer` header
171/// before upgrading the connection.
172///
173/// # Errors
174///
175/// Returns HTTP 401 if authentication fails (missing, invalid, or expired token).
176/// Returns HTTP 429 if the connection limit for this security context is reached.
177pub async fn ws_handler(
178    headers: axum::http::HeaderMap,
179    Query(params): Query<WsQueryParams>,
180    ws: WebSocketUpgrade,
181    State(state): State<RealtimeState>,
182) -> impl IntoResponse {
183    // Extract token from query param or Authorization header
184    let token = params.token.or_else(|| {
185        headers
186            .get(axum::http::header::AUTHORIZATION)
187            .and_then(|v| v.to_str().ok())
188            .and_then(|v| v.strip_prefix("Bearer "))
189            .map(str::to_owned)
190    });
191
192    let Some(token) = token else {
193        return StatusCode::UNAUTHORIZED.into_response();
194    };
195
196    // Validate token before upgrade
197    let token_info = match state.validator.validate(&token).await {
198        Ok(info) => info,
199        Err(reason) => {
200            warn!(reason = %reason, "Realtime WebSocket auth failed");
201            return StatusCode::UNAUTHORIZED.into_response();
202        },
203    };
204
205    // Check connection limit for this security context
206    let context_hash = token_info.context_hash;
207    let current = state.server.connections.count_by_context(context_hash);
208    if current >= state.server.config.max_connections_per_context {
209        return StatusCode::TOO_MANY_REQUESTS.into_response();
210    }
211
212    let server = state.server.clone();
213    ws.on_upgrade(move |socket| handle_realtime_connection(socket, server, token_info))
214        .into_response()
215}
216
217/// Handle an authenticated realtime `WebSocket` connection.
218#[allow(clippy::cognitive_complexity)] // Reason: WebSocket event loop with heartbeat, idle timeout, token expiry, and subscription handling
219async fn handle_realtime_connection(
220    socket: WebSocket,
221    server: Arc<RealtimeServer>,
222    token_info: TokenInfo,
223) {
224    let context_hash = token_info.context_hash;
225    let connection_id = uuid::Uuid::new_v4().to_string();
226    let config = &server.config;
227
228    // Register connection. Returns:
229    // - event_rx: receives change event JSON from the delivery pipeline
230    // - control_rx: fires once if the delivery pipeline detects a slow consumer
231    let conn_state = ConnectionState::new(
232        connection_id.clone(),
233        token_info.user_id.clone(),
234        token_info.context_hash,
235        token_info.expires_at,
236    );
237    let (mut event_rx, control_rx) = server.connections.insert(conn_state);
238    // Pin the close-signal receiver so it can be polled across select! loop iterations.
239    tokio::pin!(control_rx);
240
241    info!(
242        connection_id = %connection_id,
243        user_id = %token_info.user_id,
244        "Realtime WebSocket connected"
245    );
246
247    let (mut sender, mut receiver) = socket.split();
248
249    // Send connected message
250    let connected_msg = ServerMessage::Connected {
251        connection_id: connection_id.clone(),
252    };
253    if let Ok(json) = connected_msg.to_json() {
254        if sender.send(Message::Text(json.into())).await.is_err() {
255            server.connections.remove(&connection_id);
256            return;
257        }
258    }
259
260    let mut heartbeat_interval = tokio::time::interval(config.heartbeat_interval);
261    heartbeat_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
262    // Skip the immediate first tick
263    heartbeat_interval.tick().await;
264
265    let mut idle_deadline = tokio::time::Instant::now() + config.idle_timeout;
266
267    loop {
268        tokio::select! {
269            // Heartbeat tick
270            _ = heartbeat_interval.tick() => {
271                // D14: Check token expiry on each heartbeat
272                let now_ts = chrono::Utc::now().timestamp();
273                if now_ts >= token_info.expires_at {
274                    debug!(connection_id = %connection_id, "Token expired, closing connection");
275                    if let Ok(json) = ServerMessage::TokenExpired.to_json() {
276                        let _ = sender.send(Message::Text(json.into())).await;
277                    }
278                    let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
279                        code: 4401,
280                        reason: "token expired".into(),
281                    }))).await;
282                    break;
283                }
284
285                // Send ping
286                if let Ok(json) = ServerMessage::Ping.to_json() {
287                    if sender.send(Message::Text(json.into())).await.is_err() {
288                        break;
289                    }
290                }
291            }
292
293            // Idle timeout
294            () = tokio::time::sleep_until(idle_deadline) => {
295                debug!(connection_id = %connection_id, "Idle timeout, closing connection");
296                let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
297                    code: 1000,
298                    reason: "idle timeout".into(),
299                }))).await;
300                break;
301            }
302
303            // Events from delivery pipeline
304            Some(event_json) = event_rx.recv() => {
305                if sender.send(Message::Text(event_json.into())).await.is_err() {
306                    break;
307                }
308            }
309
310            // Slow-consumer close signal from the delivery pipeline
311            signal = control_rx.as_mut() => {
312                if let Ok(sig) = signal {
313                    debug!(
314                        connection_id = %connection_id,
315                        code = sig.code,
316                        reason = %sig.reason,
317                        "Slow consumer: closing connection"
318                    );
319                    let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
320                        code: sig.code,
321                        reason: sig.reason.into(),
322                    }))).await;
323                }
324                break;
325            }
326
327            // Client messages
328            msg = receiver.next() => {
329                match msg {
330                    Some(Ok(Message::Text(text))) => {
331                        // Reset idle timer on any message
332                        idle_deadline = tokio::time::Instant::now() + config.idle_timeout;
333
334                        match serde_json::from_str::<ClientMessage>(&text) {
335                            Ok(ClientMessage::Pong) => {
336                                debug!(connection_id = %connection_id, "Received pong");
337                            }
338                            Ok(ClientMessage::Subscribe { entity, event, filter }) => {
339                                let reply = handle_subscribe(
340                                    &server,
341                                    &connection_id,
342                                    context_hash,
343                                    &entity,
344                                    &event,
345                                    filter.as_deref(),
346                                );
347                                if let Ok(json) = reply.to_json() {
348                                    if sender.send(Message::Text(json.into())).await.is_err() {
349                                        break;
350                                    }
351                                }
352                            }
353                            Ok(ClientMessage::Unsubscribe { entity }) => {
354                                let _ = server.subscriptions.unsubscribe(&connection_id, &entity);
355                                let reply = ServerMessage::Unsubscribed { entity };
356                                if let Ok(json) = reply.to_json() {
357                                    if sender.send(Message::Text(json.into())).await.is_err() {
358                                        break;
359                                    }
360                                }
361                            }
362                            Err(_) => {
363                                // Unknown message, ignore
364                            }
365                        }
366                    }
367                    Some(Ok(Message::Close(_))) => {
368                        debug!(connection_id = %connection_id, "Client sent close");
369                        break;
370                    }
371                    Some(Ok(Message::Pong(_))) => {
372                        // WebSocket-level pong, reset idle timer
373                        idle_deadline = tokio::time::Instant::now() + config.idle_timeout;
374                    }
375                    Some(Err(e)) => {
376                        warn!(connection_id = %connection_id, error = %e, "WebSocket error");
377                        break;
378                    }
379                    None => break,
380                    _ => {}
381                }
382            }
383        }
384    }
385
386    // Cleanup: remove all subscriptions and connection state
387    server.subscriptions.unsubscribe_all(&connection_id);
388    server.connections.remove(&connection_id);
389    info!(connection_id = %connection_id, "Realtime WebSocket disconnected");
390}
391
392/// Handle a subscribe request from a client.
393fn handle_subscribe(
394    server: &RealtimeServer,
395    connection_id: &str,
396    context_hash: u64,
397    entity: &str,
398    event: &str,
399    filter: Option<&str>,
400) -> ServerMessage {
401    // Validate entity exists in schema
402    if !server.known_entities.is_empty() && !server.known_entities.contains(entity) {
403        return ServerMessage::Error {
404            message: format!("unknown entity: {entity}"),
405        };
406    }
407
408    // Parse event filter
409    let event_filter = if event == "*" {
410        None
411    } else {
412        match EventKind::parse(event) {
413            Ok(kind) => Some(kind),
414            Err(e) => return ServerMessage::Error { message: e },
415        }
416    };
417
418    // Parse field filters
419    let field_filters = if let Some(f) = filter {
420        match parse_filter(f) {
421            Ok(filters) => filters,
422            Err(e) => return ServerMessage::Error { message: e },
423        }
424    } else {
425        Vec::new()
426    };
427
428    let details = SubscriptionDetails {
429        event_filter,
430        field_filters,
431        security_context_hash: context_hash,
432    };
433
434    match server.subscriptions.subscribe(connection_id, entity, details) {
435        Ok(_) => ServerMessage::Subscribed {
436            entity: entity.to_owned(),
437        },
438        Err(e) => ServerMessage::Error { message: e },
439    }
440}