fraiseql_server/realtime/
server.rs1use 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#[derive(Debug, Clone)]
34pub struct RealtimeConfig {
35 pub max_connections_per_context: usize,
37 pub heartbeat_interval: Duration,
39 pub idle_timeout: Duration,
41 pub max_subscriptions_per_entity: usize,
43 pub event_channel_capacity: usize,
45 pub token_revalidation_interval: Duration,
48 pub max_consecutive_drops: usize,
51 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
70pub trait TokenValidator: Send + Sync + 'static {
85 fn validate<'a>(&'a self, token: &'a str) -> BoxFuture<'a, Result<TokenInfo, String>>;
92}
93
94#[derive(Debug, Clone)]
96pub struct TokenInfo {
97 pub user_id: String,
99 pub context_hash: u64,
101 pub expires_at: i64,
103 #[allow(clippy::struct_field_names)] pub attributes: HashMap<String, serde_json::Value>,
111 pub roles: Vec<String>,
113}
114
115impl TokenInfo {
116 #[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#[derive(Clone)]
133pub struct RealtimeState {
134 pub server: Arc<RealtimeServer>,
136 pub validator: Arc<dyn TokenValidator>,
138}
139
140pub struct RealtimeServer {
142 pub(crate) connections: Arc<ConnectionManager>,
144 pub(crate) subscriptions: Arc<SubscriptionManager>,
146 pub(crate) known_entities: HashSet<String>,
148 pub(crate) config: RealtimeConfig,
150 pub(crate) subscription_policies:
154 HashMap<String, super::subscription_policy::SubscriptionPolicy>,
155}
156
157impl RealtimeServer {
158 #[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 #[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 #[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 #[must_use]
206 pub fn active_connections(&self) -> usize {
207 self.connections.count()
208 }
209}
210
211#[derive(Debug, Deserialize)]
213pub struct WsQueryParams {
214 pub token: Option<String>,
216}
217
218pub 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 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 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 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#[allow(clippy::cognitive_complexity)] async 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 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 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 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 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_interval.tick() => {
320 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 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 () = 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 Some(event_json) = event_rx.recv() => {
354 if sender.send(Message::Text(event_json.into())).await.is_err() {
355 break;
356 }
357 }
358
359 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 msg = receiver.next() => {
378 match msg {
379 Some(Ok(Message::Text(text))) => {
380 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 }
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 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 server.subscriptions.unsubscribe_all(&connection_id);
437 server.connections.remove(&connection_id);
438 info!(connection_id = %connection_id, "Realtime WebSocket disconnected");
439}
440
441fn 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 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 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 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 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}