1use std::{
27 collections::HashMap,
28 sync::{
29 Arc,
30 atomic::{AtomicU64, Ordering},
31 },
32 time::Duration,
33};
34
35use axum::{
36 extract::{
37 State,
38 ws::{Message, WebSocket, WebSocketUpgrade},
39 },
40 http::HeaderMap,
41 response::IntoResponse,
42};
43use fraiseql_core::runtime::{
44 SubscriptionId, SubscriptionManager, SubscriptionPayload,
45 protocol::{
46 ClientMessage, ClientMessageType, CloseCode, GraphQLError, ServerMessage, SubscribePayload,
47 },
48};
49use futures::{SinkExt, StreamExt};
50use tokio::sync::broadcast;
51use tracing::{debug, error, info, warn};
52
53use crate::subscriptions::{
54 lifecycle::SubscriptionLifecycle,
55 protocol::{ProtocolCodec, WsProtocol},
56};
57
58static WS_CONNECTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
61static WS_CONNECTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
62static WS_SUBSCRIPTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
63static WS_SUBSCRIPTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
64
65#[must_use]
67pub fn subscription_metrics() -> SubscriptionMetrics {
68 SubscriptionMetrics {
69 connections_accepted: WS_CONNECTIONS_ACCEPTED.load(Ordering::Relaxed),
70 connections_rejected: WS_CONNECTIONS_REJECTED.load(Ordering::Relaxed),
71 subscriptions_accepted: WS_SUBSCRIPTIONS_ACCEPTED.load(Ordering::Relaxed),
72 subscriptions_rejected: WS_SUBSCRIPTIONS_REJECTED.load(Ordering::Relaxed),
73 }
74}
75
76#[cfg(test)]
81pub fn reset_metrics_for_test() {
82 WS_CONNECTIONS_ACCEPTED.store(0, Ordering::SeqCst);
83 WS_CONNECTIONS_REJECTED.store(0, Ordering::SeqCst);
84 WS_SUBSCRIPTIONS_ACCEPTED.store(0, Ordering::SeqCst);
85 WS_SUBSCRIPTIONS_REJECTED.store(0, Ordering::SeqCst);
86}
87
88pub struct SubscriptionMetrics {
90 pub connections_accepted: u64,
92 pub connections_rejected: u64,
94 pub subscriptions_accepted: u64,
96 pub subscriptions_rejected: u64,
98}
99
100const CONNECTION_INIT_TIMEOUT: Duration = Duration::from_secs(5);
102
103const PING_INTERVAL: Duration = Duration::from_secs(30);
105
106#[derive(Clone)]
108pub struct SubscriptionState {
109 pub manager: Arc<SubscriptionManager>,
111 pub lifecycle: Arc<dyn SubscriptionLifecycle>,
113 pub max_subscriptions_per_connection: Option<u32>,
115 pub remote_subscription_fields: Arc<HashMap<String, String>>,
120}
121
122impl SubscriptionState {
123 pub fn new(manager: Arc<SubscriptionManager>) -> Self {
125 Self {
126 manager,
127 lifecycle: Arc::new(crate::subscriptions::lifecycle::NoopLifecycle),
128 max_subscriptions_per_connection: None,
129 remote_subscription_fields: Arc::new(HashMap::new()),
130 }
131 }
132
133 #[must_use]
135 pub fn with_lifecycle(mut self, lifecycle: Arc<dyn SubscriptionLifecycle>) -> Self {
136 self.lifecycle = lifecycle;
137 self
138 }
139
140 #[must_use]
142 pub const fn with_max_subscriptions(mut self, max: Option<u32>) -> Self {
143 self.max_subscriptions_per_connection = max;
144 self
145 }
146
147 #[must_use]
151 pub fn with_remote_subscription_fields(mut self, fields: HashMap<String, String>) -> Self {
152 self.remote_subscription_fields = Arc::new(fields);
153 self
154 }
155}
156
157pub async fn subscription_handler(
164 headers: HeaderMap,
165 ws: WebSocketUpgrade,
166 State(state): State<SubscriptionState>,
167) -> impl IntoResponse {
168 let protocol_header = headers.get("sec-websocket-protocol").and_then(|v| v.to_str().ok());
169
170 let protocol = match protocol_header {
171 None => WsProtocol::GraphqlTransportWs,
172 Some(header) => {
173 if let Some(p) = WsProtocol::from_header(Some(header)) {
174 p
175 } else {
176 warn!(header = %header, "Unknown WebSocket sub-protocol requested");
177 return axum::http::StatusCode::BAD_REQUEST.into_response();
178 }
179 },
180 };
181
182 let tenant_id = super::graphql::TenantKeyResolver::resolve(None, &headers, None, false)
184 .ok()
185 .flatten();
186
187 ws.protocols([protocol.as_str()])
188 .on_upgrade(move |socket| {
189 handle_subscription_connection(socket, state, protocol, tenant_id)
190 })
191 .into_response()
192}
193
194#[allow(clippy::cognitive_complexity)] async fn handle_subscription_connection(
197 socket: WebSocket,
198 state: SubscriptionState,
199 protocol: WsProtocol,
200 tenant_id: Option<String>,
201) {
202 let connection_id = uuid::Uuid::new_v4().to_string();
203 let codec = ProtocolCodec::new(protocol);
204 info!(
205 connection_id = %connection_id,
206 protocol = %protocol.as_str(),
207 "WebSocket connection established"
208 );
209
210 let (mut sender, mut receiver) = socket.split();
211
212 let init_result = tokio::time::timeout(CONNECTION_INIT_TIMEOUT, async {
214 while let Some(msg) = receiver.next().await {
215 match msg {
216 Ok(Message::Text(text)) => {
217 if let Ok(client_msg) = codec.decode(&text) {
218 if client_msg.parsed_type() == Some(ClientMessageType::ConnectionInit) {
219 return Some(client_msg);
220 }
221 }
222 },
223 Ok(Message::Close(_)) => return None,
224 Err(e) => {
225 error!(error = %e, "WebSocket error during init");
226 return None;
227 },
228 _ => {},
229 }
230 }
231 None
232 })
233 .await;
234
235 let _init_payload = match init_result {
237 Ok(Some(msg)) => {
238 let params = msg.payload.clone().unwrap_or(serde_json::json!({}));
240 if let Err(reason) = state.lifecycle.on_connect(¶ms, &connection_id).await {
241 warn!(
242 connection_id = %connection_id,
243 reason = %reason,
244 "Lifecycle on_connect rejected connection"
245 );
246 WS_CONNECTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
247 let _ = sender
249 .send(Message::Close(Some(axum::extract::ws::CloseFrame {
250 code: 4400,
251 reason: reason.into(),
252 })))
253 .await;
254 return;
255 }
256
257 let ack = ServerMessage::connection_ack(None);
259 if let Err(send_err) = send_server_message(&codec, &mut sender, ack).await {
260 error!(connection_id = %connection_id, error = %send_err, "Failed to send connection_ack");
261 return;
262 }
263 WS_CONNECTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
264 info!(connection_id = %connection_id, "Connection initialized");
265 msg.payload
266 },
267 Ok(None) => {
268 warn!(connection_id = %connection_id, "Connection closed during init");
269 return;
270 },
271 Err(_) => {
272 warn!(connection_id = %connection_id, "Connection init timeout");
273 let _ = sender
275 .send(Message::Close(Some(axum::extract::ws::CloseFrame {
276 code: CloseCode::ConnectionInitTimeout.code(),
277 reason: CloseCode::ConnectionInitTimeout.reason().into(),
278 })))
279 .await;
280 return;
281 },
282 };
283
284 let mut active_operations: HashMap<String, SubscriptionId> = HashMap::new();
286
287 let (remote_msg_tx, mut remote_msg_rx) = tokio::sync::mpsc::channel::<ServerMessage>(64);
293
294 let mut event_receiver = state.manager.receiver();
296
297 let mut ping_interval = tokio::time::interval(PING_INTERVAL);
299 ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
300
301 loop {
321 tokio::select! {
322 msg = receiver.next() => {
323 match msg {
324 Some(Ok(Message::Text(text))) => {
325 if let Err(close_code) = handle_client_message(
326 &text,
327 &connection_id,
328 &state,
329 &codec,
330 &mut active_operations,
331 remote_msg_tx.clone(),
332 &mut sender,
333 tenant_id.as_deref(),
334 ).await {
335 let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
337 code: close_code.code(),
338 reason: close_code.reason().into(),
339 }))).await;
340 break;
341 }
342 }
343 Some(Ok(Message::Ping(data))) => {
344 let _ = sender.send(Message::Pong(data)).await;
346 }
347 Some(Ok(Message::Close(_))) => {
348 info!(connection_id = %connection_id, "Client closed connection");
349 break;
350 }
351 Some(Err(e)) => {
352 error!(connection_id = %connection_id, error = %e, "WebSocket error");
353 break;
354 }
355 None => {
356 info!(connection_id = %connection_id, "WebSocket stream ended");
357 break;
358 }
359 _ => {}
360 }
361 }
362
363 event = event_receiver.recv() => {
364 match event {
365 Ok(payload) => {
366 let tenant_matches = match (
372 tenant_id.as_deref(),
373 payload.event.tenant_id.as_deref(),
374 ) {
375 (Some(conn_tid), Some(evt_tid)) => conn_tid == evt_tid,
376 _ => true, };
378 if tenant_matches {
379 if let Some((op_id, _)) = active_operations
380 .iter()
381 .find(|(_, sub_id)| **sub_id == payload.subscription_id)
382 {
383 let msg = create_next_message(op_id, &payload);
384 if send_server_message(&codec, &mut sender, msg).await.is_err() {
385 warn!(connection_id = %connection_id, "Failed to send event");
386 break;
387 }
388 }
389 }
390 }
391 Err(broadcast::error::RecvError::Lagged(n)) => {
392 warn!(connection_id = %connection_id, lagged = n, "Event receiver lagged");
393 }
394 Err(broadcast::error::RecvError::Closed) => {
395 error!(connection_id = %connection_id, "Event channel closed");
396 break;
397 }
398 }
399 }
400
401 remote_msg = remote_msg_rx.recv() => {
402 if let Some(msg) = remote_msg {
406 if send_server_message(&codec, &mut sender, msg).await.is_err() {
407 warn!(connection_id = %connection_id, "Failed to send remote subscription message");
408 break;
409 }
410 }
411 }
412
413 _ = ping_interval.tick() => {
414 let msg = ServerMessage::ping(None);
415 if send_server_message(&codec, &mut sender, msg).await.is_err() {
416 warn!(connection_id = %connection_id, "Failed to send ping/keepalive");
417 break;
418 }
419 }
420 }
421 }
422
423 state.manager.unsubscribe_connection(&connection_id);
425 state.lifecycle.on_disconnect(&connection_id).await;
426 info!(connection_id = %connection_id, "WebSocket connection closed");
427}
428
429#[allow(clippy::cognitive_complexity)] #[allow(clippy::too_many_arguments)] async fn handle_client_message(
438 text: &str,
439 connection_id: &str,
440 state: &SubscriptionState,
441 codec: &ProtocolCodec,
442 active_operations: &mut HashMap<String, SubscriptionId>,
443 remote_msg_tx: tokio::sync::mpsc::Sender<ServerMessage>,
444 sender: &mut futures::stream::SplitSink<WebSocket, Message>,
445 tenant_id: Option<&str>,
446) -> Result<(), CloseCode> {
447 #[cfg(not(feature = "federation"))]
450 let _ = &remote_msg_tx;
451
452 let client_msg: ClientMessage = codec.decode(text).map_err(|e| {
453 warn!(error = %e, "Failed to parse client message");
454 CloseCode::ProtocolError
455 })?;
456
457 match client_msg.parsed_type() {
458 Some(ClientMessageType::Ping) => {
459 let pong = ServerMessage::pong(client_msg.payload);
460 let _ = send_server_message(codec, sender, pong).await;
462 },
463
464 Some(ClientMessageType::Pong) => {
465 debug!(connection_id = %connection_id, "Received pong");
466 },
467
468 Some(ClientMessageType::Subscribe) => {
469 let payload: SubscribePayload = client_msg.subscription_payload().ok_or_else(|| {
470 warn!("Invalid subscribe payload");
471 CloseCode::ProtocolError
472 })?;
473
474 let op_id = client_msg.id.ok_or_else(|| {
475 warn!("Subscribe message missing operation ID");
476 CloseCode::ProtocolError
477 })?;
478
479 if active_operations.contains_key(&op_id) {
481 warn!(operation_id = %op_id, "Duplicate operation ID");
482 return Err(CloseCode::SubscriberAlreadyExists);
483 }
484
485 if let Some(max) = state.max_subscriptions_per_connection {
487 if active_operations.len() >= max as usize {
488 warn!(
489 connection_id = %connection_id,
490 active = active_operations.len(),
491 max = max,
492 "Subscription limit reached"
493 );
494 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
495 let error = ServerMessage::error(
496 &op_id,
497 vec![GraphQLError::with_code(
498 format!("Maximum subscriptions per connection ({max}) reached"),
499 "SUBSCRIPTION_LIMIT_REACHED",
500 )],
501 );
502 if let Err(e) = send_server_message(codec, sender, error).await {
503 debug!(connection_id = %connection_id, error = %e, "Could not send subscription limit error to client");
504 }
505 return Ok(());
506 }
507 }
508
509 let Some(subscription_name) = extract_subscription_name(&payload.query) else {
511 let error = ServerMessage::error(
512 &op_id,
513 vec![GraphQLError::with_code(
514 "Could not parse subscription query",
515 "PARSE_ERROR",
516 )],
517 );
518 if let Err(e) = send_server_message(codec, sender, error).await {
519 debug!(connection_id = %connection_id, error = %e, "Could not send parse error to client");
520 }
521 return Ok(());
522 };
523
524 let variables_value = serde_json::to_value(&payload.variables)
527 .expect("HashMap<String, serde_json::Value> serialization is infallible");
528 if let Err(reason) = state
529 .lifecycle
530 .on_subscribe(&subscription_name, &variables_value, connection_id)
531 .await
532 {
533 warn!(
534 connection_id = %connection_id,
535 subscription = %subscription_name,
536 reason = %reason,
537 "Lifecycle on_subscribe rejected subscription"
538 );
539 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
540 let error = ServerMessage::error(
541 &op_id,
542 vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REJECTED")],
543 );
544 if let Err(e) = send_server_message(codec, sender, error).await {
545 debug!(connection_id = %connection_id, error = %e, "Could not send subscription rejection to client");
546 }
547 return Ok(());
548 }
549
550 #[cfg(feature = "federation")]
552 if let Some(subgraph_url) = state.remote_subscription_fields.get(&subscription_name) {
553 use fraiseql_federation::subscription_forwarder::{
554 ForwardedEvent, SubscriptionForwarder,
555 };
556
557 match SubscriptionForwarder::new(subgraph_url) {
558 Ok(forwarder) => {
559 let (event_tx, mut event_rx) =
561 tokio::sync::mpsc::channel::<ForwardedEvent>(32);
562
563 let fwd_op = op_id.clone();
565 let fwd_query = payload.query.clone();
566 let fwd_vars = variables_value.clone();
567 tokio::spawn(async move {
568 if let Err(e) =
569 forwarder.forward(&fwd_op, &fwd_query, fwd_vars, event_tx).await
570 {
571 warn!(error = %e, "Remote subscription forwarder failed");
572 }
573 });
574
575 let relay_op = op_id.clone();
577 let relay_tx = remote_msg_tx.clone();
578 tokio::spawn(async move {
579 while let Some(event) = event_rx.recv().await {
580 let server_msg = match event {
581 ForwardedEvent::Next(data) => {
582 ServerMessage::next(&relay_op, data)
583 },
584 ForwardedEvent::Error(errors) => {
585 let errors_vec = errors.as_array().map_or_else(
586 || {
587 vec![GraphQLError::with_code(
588 errors.to_string(),
589 "REMOTE_ERROR",
590 )]
591 },
592 |arr| {
593 arr.iter()
594 .map(|e| {
595 GraphQLError::with_code(
596 e.get("message")
597 .and_then(|v| v.as_str())
598 .unwrap_or("Remote subgraph error"),
599 "REMOTE_ERROR",
600 )
601 })
602 .collect()
603 },
604 );
605 ServerMessage::error(&relay_op, errors_vec)
606 },
607 ForwardedEvent::Complete => ServerMessage::complete(&relay_op),
608 };
609 if relay_tx.send(server_msg).await.is_err() {
610 break; }
612 }
613 });
614
615 WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
616 info!(
617 connection_id = %connection_id,
618 operation_id = %op_id,
619 subscription = %subscription_name,
620 "Subscription forwarded to remote subgraph"
621 );
622 return Ok(());
623 },
624 Err(e) => {
625 let error = ServerMessage::error(
626 &op_id,
627 vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
628 );
629 if let Err(send_err) = send_server_message(codec, sender, error).await {
630 debug!(connection_id = %connection_id, error = %send_err, "Could not send forwarding error to client");
631 }
632 return Ok(());
633 },
634 }
635 }
636
637 if let Some(server_tid) = tenant_id {
639 if let Some(client_tid) = variables_value.get("tenant_id").and_then(|v| v.as_str())
640 {
641 if client_tid != server_tid {
642 let error = ServerMessage::error(
643 &op_id,
644 vec![GraphQLError::with_code(
645 format!(
646 "Tenant mismatch: client provided '{client_tid}', server resolved '{server_tid}'"
647 ),
648 "TENANT_MISMATCH",
649 )],
650 );
651 if let Err(send_err) = send_server_message(codec, sender, error).await {
652 debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant mismatch error to client");
653 }
654 return Ok(());
655 }
656 }
657 }
658
659 let mut context = serde_json::json!({});
661 if let Some(tid) = tenant_id {
662 context["tenant_id"] = serde_json::Value::String(tid.to_string());
663 }
664
665 match state.manager.subscribe(
667 &subscription_name,
668 context,
669 variables_value,
670 connection_id,
671 ) {
672 Ok(sub_id) => {
673 active_operations.insert(op_id.clone(), sub_id);
674 WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
675 info!(
676 connection_id = %connection_id,
677 operation_id = %op_id,
678 subscription = %subscription_name,
679 "Subscription started"
680 );
681 },
682 Err(e) => {
683 let error = ServerMessage::error(
684 &op_id,
685 vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
686 );
687 if let Err(send_err) = send_server_message(codec, sender, error).await {
688 debug!(connection_id = %connection_id, error = %send_err, "Could not send subscription error to client");
689 }
690 },
691 }
692 },
693
694 Some(ClientMessageType::Complete) => {
695 let op_id = client_msg.id.ok_or_else(|| {
696 warn!("Complete message missing operation ID");
697 CloseCode::ProtocolError
698 })?;
699
700 if let Some(sub_id) = active_operations.remove(&op_id) {
701 if let Err(e) = state.manager.unsubscribe(sub_id) {
702 warn!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe; subscription may be leaked");
703 }
704 state.lifecycle.on_unsubscribe(&op_id, connection_id).await;
705 info!(
706 connection_id = %connection_id,
707 operation_id = %op_id,
708 "Subscription completed"
709 );
710 }
711 },
712
713 Some(ClientMessageType::ConnectionInit) => {
714 warn!(connection_id = %connection_id, "Duplicate connection_init");
715 return Err(CloseCode::TooManyInitRequests);
716 },
717
718 None => {
719 warn!(message_type = %client_msg.message_type, "Unknown message type");
720 },
721 _ => {
723 warn!(message_type = %client_msg.message_type, "Unrecognized message type");
724 },
725 }
726
727 Ok(())
728}
729
730async fn send_server_message(
732 codec: &ProtocolCodec,
733 sender: &mut futures::stream::SplitSink<WebSocket, Message>,
734 msg: ServerMessage,
735) -> Result<(), String> {
736 match codec.encode(&msg) {
737 Ok(Some(json)) => sender.send(Message::Text(json.into())).await.map_err(|e| e.to_string()),
738 Ok(None) => Ok(()), Err(e) => Err(e.to_string()),
740 }
741}
742
743fn create_next_message(operation_id: &str, payload: &SubscriptionPayload) -> ServerMessage {
745 let data = serde_json::json!({
746 payload.subscription_name.clone(): payload.data
747 });
748 ServerMessage::next(operation_id, data)
749}
750
751pub(crate) fn extract_subscription_name(query: &str) -> Option<String> {
753 let query = query.trim();
754
755 let sub_idx = query.find("subscription")?;
756 let after_sub = &query[sub_idx + "subscription".len()..];
757
758 let brace_idx = after_sub.find('{')?;
759 let after_brace = after_sub[brace_idx + 1..].trim_start();
760
761 let name_end = after_brace
762 .find(|c: char| !c.is_alphanumeric() && c != '_')
763 .unwrap_or(after_brace.len());
764
765 if name_end == 0 {
766 return None;
767 }
768
769 Some(after_brace[..name_end].to_string())
770}