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::{
44 runtime::{
45 SubscriptionId, SubscriptionManager, SubscriptionPayload,
46 protocol::{
47 ClientMessage, ClientMessageType, CloseCode, GraphQLError, ServerMessage,
48 SubscribePayload,
49 },
50 },
51 security::{Authorizer, OperationKind, SecurityContext, authorizer::enforce_authz},
52};
53use futures::{SinkExt, StreamExt};
54use tokio::sync::broadcast;
55use tracing::{debug, error, info, warn};
56
57use crate::{
58 extractors::OptionalSecurityContext,
59 routes::graphql::{DomainRegistry, TenantStatusSource},
60 subscriptions::{
61 lifecycle::SubscriptionLifecycle,
62 protocol::{ProtocolCodec, WsProtocol},
63 },
64};
65
66static WS_CONNECTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
69static WS_CONNECTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
70static WS_SUBSCRIPTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
71static WS_SUBSCRIPTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
72
73#[must_use]
75pub fn subscription_metrics() -> SubscriptionMetrics {
76 SubscriptionMetrics {
77 connections_accepted: WS_CONNECTIONS_ACCEPTED.load(Ordering::Relaxed),
78 connections_rejected: WS_CONNECTIONS_REJECTED.load(Ordering::Relaxed),
79 subscriptions_accepted: WS_SUBSCRIPTIONS_ACCEPTED.load(Ordering::Relaxed),
80 subscriptions_rejected: WS_SUBSCRIPTIONS_REJECTED.load(Ordering::Relaxed),
81 }
82}
83
84#[cfg(test)]
89pub fn reset_metrics_for_test() {
90 WS_CONNECTIONS_ACCEPTED.store(0, Ordering::SeqCst);
91 WS_CONNECTIONS_REJECTED.store(0, Ordering::SeqCst);
92 WS_SUBSCRIPTIONS_ACCEPTED.store(0, Ordering::SeqCst);
93 WS_SUBSCRIPTIONS_REJECTED.store(0, Ordering::SeqCst);
94}
95
96pub struct SubscriptionMetrics {
98 pub connections_accepted: u64,
100 pub connections_rejected: u64,
102 pub subscriptions_accepted: u64,
104 pub subscriptions_rejected: u64,
106}
107
108const CONNECTION_INIT_TIMEOUT: Duration = Duration::from_secs(5);
110
111const PING_INTERVAL: Duration = Duration::from_secs(30);
113
114#[derive(Clone)]
116pub struct SubscriptionState {
117 pub manager: Arc<SubscriptionManager>,
119 pub lifecycle: Arc<dyn SubscriptionLifecycle>,
121 pub max_subscriptions_per_connection: Option<u32>,
123 pub remote_subscription_fields: Arc<HashMap<String, String>>,
128 pub domain_registry: Option<Arc<DomainRegistry>>,
131 pub strict_tenant_validation: bool,
135 pub authorizer: Option<Arc<dyn Authorizer>>,
140 pub tenant_status_source: Option<Arc<dyn TenantStatusSource>>,
145}
146
147impl SubscriptionState {
148 pub fn new(manager: Arc<SubscriptionManager>) -> Self {
150 Self {
151 manager,
152 lifecycle: Arc::new(crate::subscriptions::lifecycle::NoopLifecycle),
153 max_subscriptions_per_connection: None,
154 remote_subscription_fields: Arc::new(HashMap::new()),
155 domain_registry: None,
156 strict_tenant_validation: false,
157 authorizer: None,
158 tenant_status_source: None,
159 }
160 }
161
162 #[must_use]
166 pub fn with_tenant_status_source(
167 mut self,
168 source: Option<Arc<dyn TenantStatusSource>>,
169 ) -> Self {
170 self.tenant_status_source = source;
171 self
172 }
173
174 #[must_use]
179 pub fn with_authorizer(mut self, authorizer: Option<Arc<dyn Authorizer>>) -> Self {
180 self.authorizer = authorizer;
181 self
182 }
183
184 #[must_use]
191 pub fn with_tenant_context(
192 mut self,
193 domain_registry: Arc<DomainRegistry>,
194 strict_tenant_validation: bool,
195 ) -> Self {
196 self.domain_registry = Some(domain_registry);
197 self.strict_tenant_validation = strict_tenant_validation;
198 self
199 }
200
201 #[must_use]
203 pub fn with_lifecycle(mut self, lifecycle: Arc<dyn SubscriptionLifecycle>) -> Self {
204 self.lifecycle = lifecycle;
205 self
206 }
207
208 #[must_use]
210 pub const fn with_max_subscriptions(mut self, max: Option<u32>) -> Self {
211 self.max_subscriptions_per_connection = max;
212 self
213 }
214
215 #[must_use]
219 pub fn with_remote_subscription_fields(mut self, fields: HashMap<String, String>) -> Self {
220 self.remote_subscription_fields = Arc::new(fields);
221 self
222 }
223}
224
225pub async fn subscription_handler(
232 headers: HeaderMap,
233 OptionalSecurityContext(security_context): OptionalSecurityContext,
234 ws: WebSocketUpgrade,
235 State(state): State<SubscriptionState>,
236) -> impl IntoResponse {
237 let protocol_header = headers.get("sec-websocket-protocol").and_then(|v| v.to_str().ok());
238
239 let protocol = match protocol_header {
240 None => WsProtocol::GraphqlTransportWs,
241 Some(header) => {
242 if let Some(p) = WsProtocol::from_header(Some(header)) {
243 p
244 } else {
245 warn!(header = %header, "Unknown WebSocket sub-protocol requested");
246 return axum::http::StatusCode::BAD_REQUEST.into_response();
247 }
248 },
249 };
250
251 let tenant_id = match resolve_subscription_tenant(security_context.as_ref(), &headers, &state) {
258 Ok(tenant_id) => tenant_id,
259 Err(e) => {
260 warn!(error = %e, "Subscription tenant resolution rejected the upgrade");
261 return axum::http::StatusCode::BAD_REQUEST.into_response();
262 },
263 };
264
265 ws.protocols([protocol.as_str()])
266 .on_upgrade(move |socket| {
267 handle_subscription_connection(socket, state, protocol, tenant_id, security_context)
268 })
269 .into_response()
270}
271
272fn resolve_subscription_tenant(
282 security_context: Option<&SecurityContext>,
283 headers: &HeaderMap,
284 state: &SubscriptionState,
285) -> fraiseql_error::Result<Option<String>> {
286 super::graphql::TenantKeyResolver::resolve(
287 security_context,
288 headers,
289 state.domain_registry.as_deref(),
290 state.strict_tenant_validation,
291 )
292}
293
294#[allow(clippy::cognitive_complexity)] async fn handle_subscription_connection(
297 socket: WebSocket,
298 state: SubscriptionState,
299 protocol: WsProtocol,
300 tenant_id: Option<String>,
301 principal: Option<SecurityContext>,
302) {
303 let connection_id = uuid::Uuid::new_v4().to_string();
304 let codec = ProtocolCodec::new(protocol);
305 info!(
306 connection_id = %connection_id,
307 protocol = %protocol.as_str(),
308 "WebSocket connection established"
309 );
310
311 let (mut sender, mut receiver) = socket.split();
312
313 let init_result = tokio::time::timeout(CONNECTION_INIT_TIMEOUT, async {
315 while let Some(msg) = receiver.next().await {
316 match msg {
317 Ok(Message::Text(text)) => {
318 if let Ok(client_msg) = codec.decode(&text) {
319 if client_msg.parsed_type() == Some(ClientMessageType::ConnectionInit) {
320 return Some(client_msg);
321 }
322 }
323 },
324 Ok(Message::Close(_)) => return None,
325 Err(e) => {
326 error!(error = %e, "WebSocket error during init");
327 return None;
328 },
329 _ => {},
330 }
331 }
332 None
333 })
334 .await;
335
336 let _init_payload = match init_result {
338 Ok(Some(msg)) => {
339 let params = msg.payload.clone().unwrap_or(serde_json::json!({}));
341 if let Err(reason) = state.lifecycle.on_connect(¶ms, &connection_id).await {
342 warn!(
343 connection_id = %connection_id,
344 reason = %reason,
345 "Lifecycle on_connect rejected connection"
346 );
347 WS_CONNECTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
348 let _ = sender
350 .send(Message::Close(Some(axum::extract::ws::CloseFrame {
351 code: 4400,
352 reason: reason.into(),
353 })))
354 .await;
355 return;
356 }
357
358 let ack = ServerMessage::connection_ack(None);
360 if let Err(send_err) = send_server_message(&codec, &mut sender, ack).await {
361 error!(connection_id = %connection_id, error = %send_err, "Failed to send connection_ack");
362 return;
363 }
364 WS_CONNECTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
365 info!(connection_id = %connection_id, "Connection initialized");
366 msg.payload
367 },
368 Ok(None) => {
369 warn!(connection_id = %connection_id, "Connection closed during init");
370 return;
371 },
372 Err(_) => {
373 warn!(connection_id = %connection_id, "Connection init timeout");
374 let _ = sender
376 .send(Message::Close(Some(axum::extract::ws::CloseFrame {
377 code: CloseCode::ConnectionInitTimeout.code(),
378 reason: CloseCode::ConnectionInitTimeout.reason().into(),
379 })))
380 .await;
381 return;
382 },
383 };
384
385 let mut active_operations: HashMap<String, SubscriptionId> = HashMap::new();
387
388 let (remote_msg_tx, mut remote_msg_rx) = tokio::sync::mpsc::channel::<ServerMessage>(64);
394
395 let mut event_receiver = state.manager.receiver();
397
398 let mut ping_interval = tokio::time::interval(PING_INTERVAL);
400 ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
401
402 loop {
422 tokio::select! {
423 msg = receiver.next() => {
424 match msg {
425 Some(Ok(Message::Text(text))) => {
426 if let Err(close_code) = handle_client_message(
427 &text,
428 &connection_id,
429 &state,
430 &codec,
431 &mut active_operations,
432 remote_msg_tx.clone(),
433 &mut sender,
434 tenant_id.as_deref(),
435 principal.as_ref(),
436 ).await {
437 let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
439 code: close_code.code(),
440 reason: close_code.reason().into(),
441 }))).await;
442 break;
443 }
444 }
445 Some(Ok(Message::Ping(data))) => {
446 let _ = sender.send(Message::Pong(data)).await;
448 }
449 Some(Ok(Message::Close(_))) => {
450 info!(connection_id = %connection_id, "Client closed connection");
451 break;
452 }
453 Some(Err(e)) => {
454 error!(connection_id = %connection_id, error = %e, "WebSocket error");
455 break;
456 }
457 None => {
458 info!(connection_id = %connection_id, "WebSocket stream ended");
459 break;
460 }
461 _ => {}
462 }
463 }
464
465 event = event_receiver.recv() => {
466 match event {
467 Ok(payload) => {
468 let tenant_matches = match (
474 tenant_id.as_deref(),
475 payload.event.tenant_id.as_deref(),
476 ) {
477 (Some(conn_tid), Some(evt_tid)) => conn_tid == evt_tid,
478 _ => true, };
480 let tenant_active = match (
484 tenant_id.as_deref(),
485 state.tenant_status_source.as_ref(),
486 ) {
487 (Some(tid), Some(src)) => !src.is_suspended(tid),
488 _ => true,
489 };
490 if tenant_matches && tenant_active {
491 if let Some((op_id, _)) = active_operations
492 .iter()
493 .find(|(_, sub_id)| **sub_id == payload.subscription_id)
494 {
495 let msg = create_next_message(op_id, &payload);
496 if send_server_message(&codec, &mut sender, msg).await.is_err() {
497 warn!(connection_id = %connection_id, "Failed to send event");
498 break;
499 }
500 }
501 }
502 }
503 Err(broadcast::error::RecvError::Lagged(n)) => {
504 warn!(connection_id = %connection_id, lagged = n, "Event receiver lagged");
505 }
506 Err(broadcast::error::RecvError::Closed) => {
507 error!(connection_id = %connection_id, "Event channel closed");
508 break;
509 }
510 }
511 }
512
513 remote_msg = remote_msg_rx.recv() => {
514 if let Some(msg) = remote_msg {
518 if send_server_message(&codec, &mut sender, msg).await.is_err() {
519 warn!(connection_id = %connection_id, "Failed to send remote subscription message");
520 break;
521 }
522 }
523 }
524
525 _ = ping_interval.tick() => {
526 let msg = ServerMessage::ping(None);
527 if send_server_message(&codec, &mut sender, msg).await.is_err() {
528 warn!(connection_id = %connection_id, "Failed to send ping/keepalive");
529 break;
530 }
531 }
532 }
533 }
534
535 state.manager.unsubscribe_connection(&connection_id);
537 state.lifecycle.on_disconnect(&connection_id).await;
538 info!(connection_id = %connection_id, "WebSocket connection closed");
539}
540
541#[allow(clippy::cognitive_complexity)] #[allow(clippy::too_many_arguments)] async fn handle_client_message(
550 text: &str,
551 connection_id: &str,
552 state: &SubscriptionState,
553 codec: &ProtocolCodec,
554 active_operations: &mut HashMap<String, SubscriptionId>,
555 remote_msg_tx: tokio::sync::mpsc::Sender<ServerMessage>,
556 sender: &mut futures::stream::SplitSink<WebSocket, Message>,
557 tenant_id: Option<&str>,
558 principal: Option<&SecurityContext>,
559) -> Result<(), CloseCode> {
560 #[cfg(not(feature = "federation"))]
563 let _ = &remote_msg_tx;
564
565 let client_msg: ClientMessage = codec.decode(text).map_err(|e| {
566 warn!(error = %e, "Failed to parse client message");
567 CloseCode::ProtocolError
568 })?;
569
570 match client_msg.parsed_type() {
571 Some(ClientMessageType::Ping) => {
572 let pong = ServerMessage::pong(client_msg.payload);
573 let _ = send_server_message(codec, sender, pong).await;
575 },
576
577 Some(ClientMessageType::Pong) => {
578 debug!(connection_id = %connection_id, "Received pong");
579 },
580
581 Some(ClientMessageType::Subscribe) => {
582 let payload: SubscribePayload = client_msg.subscription_payload().ok_or_else(|| {
583 warn!("Invalid subscribe payload");
584 CloseCode::ProtocolError
585 })?;
586
587 let op_id = client_msg.id.ok_or_else(|| {
588 warn!("Subscribe message missing operation ID");
589 CloseCode::ProtocolError
590 })?;
591
592 if active_operations.contains_key(&op_id) {
594 warn!(operation_id = %op_id, "Duplicate operation ID");
595 return Err(CloseCode::SubscriberAlreadyExists);
596 }
597
598 if let Some(max) = state.max_subscriptions_per_connection {
600 if active_operations.len() >= max as usize {
601 warn!(
602 connection_id = %connection_id,
603 active = active_operations.len(),
604 max = max,
605 "Subscription limit reached"
606 );
607 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
608 let error = ServerMessage::error(
609 &op_id,
610 vec![GraphQLError::with_code(
611 format!("Maximum subscriptions per connection ({max}) reached"),
612 "SUBSCRIPTION_LIMIT_REACHED",
613 )],
614 );
615 if let Err(e) = send_server_message(codec, sender, error).await {
616 debug!(connection_id = %connection_id, error = %e, "Could not send subscription limit error to client");
617 }
618 return Ok(());
619 }
620 }
621
622 let Some(subscription_name) = extract_subscription_name(&payload.query) else {
624 let error = ServerMessage::error(
625 &op_id,
626 vec![GraphQLError::with_code(
627 "Could not parse subscription query",
628 "PARSE_ERROR",
629 )],
630 );
631 if let Err(e) = send_server_message(codec, sender, error).await {
632 debug!(connection_id = %connection_id, error = %e, "Could not send parse error to client");
633 }
634 return Ok(());
635 };
636
637 let variables_value = serde_json::to_value(&payload.variables)
640 .expect("HashMap<String, serde_json::Value> serialization is infallible");
641
642 if let Some(authorizer) = state.authorizer.as_ref() {
648 let ops = [(OperationKind::Subscription, subscription_name.clone())];
649 if let Err(err) =
650 enforce_authz(authorizer.as_ref(), principal, &ops, Some(&variables_value))
651 {
652 warn!(
653 connection_id = %connection_id,
654 subscription = %subscription_name,
655 "Operation authorizer denied the subscription"
656 );
657 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
658 let error = ServerMessage::error(
659 &op_id,
660 vec![GraphQLError::with_code(err.to_string(), "FORBIDDEN")],
661 );
662 if let Err(e) = send_server_message(codec, sender, error).await {
663 debug!(connection_id = %connection_id, error = %e, "Could not send authorization denial to client");
664 }
665 return Ok(());
666 }
667 }
668
669 if let Err(reason) = state
670 .lifecycle
671 .on_subscribe(&subscription_name, &variables_value, connection_id)
672 .await
673 {
674 warn!(
675 connection_id = %connection_id,
676 subscription = %subscription_name,
677 reason = %reason,
678 "Lifecycle on_subscribe rejected subscription"
679 );
680 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
681 let error = ServerMessage::error(
682 &op_id,
683 vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REJECTED")],
684 );
685 if let Err(e) = send_server_message(codec, sender, error).await {
686 debug!(connection_id = %connection_id, error = %e, "Could not send subscription rejection to client");
687 }
688 return Ok(());
689 }
690
691 #[cfg(feature = "federation")]
693 if let Some(subgraph_url) = state.remote_subscription_fields.get(&subscription_name) {
694 use fraiseql_federation::subscription_forwarder::{
695 ForwardedEvent, SubscriptionForwarder,
696 };
697
698 match SubscriptionForwarder::new(subgraph_url) {
699 Ok(forwarder) => {
700 let (event_tx, mut event_rx) =
702 tokio::sync::mpsc::channel::<ForwardedEvent>(32);
703
704 let fwd_op = op_id.clone();
706 let fwd_query = payload.query.clone();
707 let fwd_vars = variables_value.clone();
708 tokio::spawn(async move {
709 if let Err(e) =
710 forwarder.forward(&fwd_op, &fwd_query, fwd_vars, event_tx).await
711 {
712 warn!(error = %e, "Remote subscription forwarder failed");
713 }
714 });
715
716 let relay_op = op_id.clone();
718 let relay_tx = remote_msg_tx.clone();
719 tokio::spawn(async move {
720 while let Some(event) = event_rx.recv().await {
721 let server_msg = match event {
722 ForwardedEvent::Next(data) => {
723 ServerMessage::next(&relay_op, data)
724 },
725 ForwardedEvent::Error(errors) => {
726 let errors_vec = errors.as_array().map_or_else(
727 || {
728 vec![GraphQLError::with_code(
729 errors.to_string(),
730 "REMOTE_ERROR",
731 )]
732 },
733 |arr| {
734 arr.iter()
735 .map(|e| {
736 GraphQLError::with_code(
737 e.get("message")
738 .and_then(|v| v.as_str())
739 .unwrap_or("Remote subgraph error"),
740 "REMOTE_ERROR",
741 )
742 })
743 .collect()
744 },
745 );
746 ServerMessage::error(&relay_op, errors_vec)
747 },
748 ForwardedEvent::Complete => ServerMessage::complete(&relay_op),
749 };
750 if relay_tx.send(server_msg).await.is_err() {
751 break; }
753 }
754 });
755
756 WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
757 info!(
758 connection_id = %connection_id,
759 operation_id = %op_id,
760 subscription = %subscription_name,
761 "Subscription forwarded to remote subgraph"
762 );
763 return Ok(());
764 },
765 Err(e) => {
766 let error = ServerMessage::error(
767 &op_id,
768 vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
769 );
770 if let Err(send_err) = send_server_message(codec, sender, error).await {
771 debug!(connection_id = %connection_id, error = %send_err, "Could not send forwarding error to client");
772 }
773 return Ok(());
774 },
775 }
776 }
777
778 if let Some(server_tid) = tenant_id {
780 if let Some(client_tid) = variables_value.get("tenant_id").and_then(|v| v.as_str())
781 {
782 if client_tid != server_tid {
783 let error = ServerMessage::error(
784 &op_id,
785 vec![GraphQLError::with_code(
786 format!(
787 "Tenant mismatch: client provided '{client_tid}', server resolved '{server_tid}'"
788 ),
789 "TENANT_MISMATCH",
790 )],
791 );
792 if let Err(send_err) = send_server_message(codec, sender, error).await {
793 debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant mismatch error to client");
794 }
795 return Ok(());
796 }
797 }
798 }
799
800 if let (Some(tid), Some(src)) = (tenant_id, state.tenant_status_source.as_ref()) {
803 if src.is_suspended(tid) {
804 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
805 let error = ServerMessage::error(
806 &op_id,
807 vec![GraphQLError::with_code(
808 format!("Tenant '{tid}' is suspended"),
809 "TENANT_SUSPENDED",
810 )],
811 );
812 if let Err(send_err) = send_server_message(codec, sender, error).await {
813 debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant-suspended error to client");
814 }
815 return Ok(());
816 }
817 }
818
819 let mut context = serde_json::json!({});
821 if let Some(tid) = tenant_id {
822 context["tenant_id"] = serde_json::Value::String(tid.to_string());
823 }
824
825 match state.manager.subscribe(
827 &subscription_name,
828 context,
829 variables_value,
830 connection_id,
831 ) {
832 Ok(sub_id) => {
833 active_operations.insert(op_id.clone(), sub_id);
834 WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
835 info!(
836 connection_id = %connection_id,
837 operation_id = %op_id,
838 subscription = %subscription_name,
839 "Subscription started"
840 );
841 },
842 Err(e) => {
843 let error = ServerMessage::error(
844 &op_id,
845 vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
846 );
847 if let Err(send_err) = send_server_message(codec, sender, error).await {
848 debug!(connection_id = %connection_id, error = %send_err, "Could not send subscription error to client");
849 }
850 },
851 }
852 },
853
854 Some(ClientMessageType::Complete) => {
855 let op_id = client_msg.id.ok_or_else(|| {
856 warn!("Complete message missing operation ID");
857 CloseCode::ProtocolError
858 })?;
859
860 if let Some(sub_id) = active_operations.remove(&op_id) {
861 if let Err(e) = state.manager.unsubscribe(sub_id) {
862 warn!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe; subscription may be leaked");
863 }
864 state.lifecycle.on_unsubscribe(&op_id, connection_id).await;
865 info!(
866 connection_id = %connection_id,
867 operation_id = %op_id,
868 "Subscription completed"
869 );
870 }
871 },
872
873 Some(ClientMessageType::ConnectionInit) => {
874 warn!(connection_id = %connection_id, "Duplicate connection_init");
875 return Err(CloseCode::TooManyInitRequests);
876 },
877
878 None => {
879 warn!(message_type = %client_msg.message_type, "Unknown message type");
880 },
881 _ => {
883 warn!(message_type = %client_msg.message_type, "Unrecognized message type");
884 },
885 }
886
887 Ok(())
888}
889
890async fn send_server_message(
892 codec: &ProtocolCodec,
893 sender: &mut futures::stream::SplitSink<WebSocket, Message>,
894 msg: ServerMessage,
895) -> Result<(), String> {
896 match codec.encode(&msg) {
897 Ok(Some(json)) => sender.send(Message::Text(json.into())).await.map_err(|e| e.to_string()),
898 Ok(None) => Ok(()), Err(e) => Err(e.to_string()),
900 }
901}
902
903fn create_next_message(operation_id: &str, payload: &SubscriptionPayload) -> ServerMessage {
910 let data = serde_json::json!({
911 payload.subscription_name.clone(): payload.data
912 });
913 match &payload.event.change_spine {
914 Some(envelope) => {
915 let extensions = serde_json::json!({ "changeSpine": envelope });
916 ServerMessage::next_with_extensions(operation_id, data, extensions)
917 },
918 None => ServerMessage::next(operation_id, data),
919 }
920}
921
922pub(crate) fn extract_subscription_name(query: &str) -> Option<String> {
924 let query = query.trim();
925
926 let sub_idx = query.find("subscription")?;
927 let after_sub = &query[sub_idx + "subscription".len()..];
928
929 let brace_idx = after_sub.find('{')?;
930 let after_brace = after_sub[brace_idx + 1..].trim_start();
931
932 let name_end = after_brace
933 .find(|c: char| !c.is_alphanumeric() && c != '_')
934 .unwrap_or(after_brace.len());
935
936 if name_end == 0 {
937 return None;
938 }
939
940 Some(after_brace[..name_end].to_string())
941}
942
943#[cfg(test)]
944mod tests;