fraiseql_server/realtime/
server.rs1use 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#[derive(Debug, Clone)]
29pub struct RealtimeConfig {
30 pub max_connections_per_context: usize,
32 pub heartbeat_interval: Duration,
34 pub idle_timeout: Duration,
36 pub max_subscriptions_per_entity: usize,
38 pub event_channel_capacity: usize,
40 pub token_revalidation_interval: Duration,
43 pub max_consecutive_drops: usize,
46 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
65pub trait TokenValidator: Send + Sync + 'static {
80 fn validate<'a>(&'a self, token: &'a str) -> BoxFuture<'a, Result<TokenInfo, String>>;
87}
88
89#[derive(Debug, Clone)]
91pub struct TokenInfo {
92 pub user_id: String,
94 pub context_hash: u64,
96 pub expires_at: i64,
98}
99
100#[derive(Clone)]
102pub struct RealtimeState {
103 pub server: Arc<RealtimeServer>,
105 pub validator: Arc<dyn TokenValidator>,
107}
108
109pub struct RealtimeServer {
111 pub(crate) connections: Arc<ConnectionManager>,
113 pub(crate) subscriptions: Arc<SubscriptionManager>,
115 pub(crate) known_entities: HashSet<String>,
117 pub(crate) config: RealtimeConfig,
119}
120
121impl RealtimeServer {
122 #[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 #[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 #[must_use]
156 pub fn active_connections(&self) -> usize {
157 self.connections.count()
158 }
159}
160
161#[derive(Debug, Deserialize)]
163pub struct WsQueryParams {
164 pub token: Option<String>,
166}
167
168pub 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 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 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 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#[allow(clippy::cognitive_complexity)] async 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 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 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 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 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_interval.tick() => {
271 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 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 () = 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 Some(event_json) = event_rx.recv() => {
305 if sender.send(Message::Text(event_json.into())).await.is_err() {
306 break;
307 }
308 }
309
310 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 msg = receiver.next() => {
329 match msg {
330 Some(Ok(Message::Text(text))) => {
331 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 }
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 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 server.subscriptions.unsubscribe_all(&connection_id);
388 server.connections.remove(&connection_id);
389 info!(connection_id = %connection_id, "Realtime WebSocket disconnected");
390}
391
392fn 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 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 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 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}