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