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 schema::{CompiledSchema, OwnerCondition, SubscriptionPolicy},
52 security::{Authorizer, OperationKind, SecurityContext, authorizer::enforce_authz},
53};
54use futures::{SinkExt, StreamExt};
55use tokio::sync::broadcast;
56use tracing::{debug, error, info, warn};
57
58use crate::{
59 extractors::OptionalSecurityContext,
60 routes::graphql::{DomainRegistry, TenantStatusSource},
61 subscriptions::{
62 lifecycle::SubscriptionLifecycle,
63 protocol::{ProtocolCodec, WsProtocol},
64 },
65};
66
67static WS_CONNECTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
70static WS_CONNECTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
71static WS_SUBSCRIPTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
72static WS_SUBSCRIPTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
73
74#[must_use]
76pub fn subscription_metrics() -> SubscriptionMetrics {
77 SubscriptionMetrics {
78 connections_accepted: WS_CONNECTIONS_ACCEPTED.load(Ordering::Relaxed),
79 connections_rejected: WS_CONNECTIONS_REJECTED.load(Ordering::Relaxed),
80 subscriptions_accepted: WS_SUBSCRIPTIONS_ACCEPTED.load(Ordering::Relaxed),
81 subscriptions_rejected: WS_SUBSCRIPTIONS_REJECTED.load(Ordering::Relaxed),
82 }
83}
84
85#[cfg(test)]
90pub fn reset_metrics_for_test() {
91 WS_CONNECTIONS_ACCEPTED.store(0, Ordering::SeqCst);
92 WS_CONNECTIONS_REJECTED.store(0, Ordering::SeqCst);
93 WS_SUBSCRIPTIONS_ACCEPTED.store(0, Ordering::SeqCst);
94 WS_SUBSCRIPTIONS_REJECTED.store(0, Ordering::SeqCst);
95}
96
97pub struct SubscriptionMetrics {
99 pub connections_accepted: u64,
101 pub connections_rejected: u64,
103 pub subscriptions_accepted: u64,
105 pub subscriptions_rejected: u64,
107}
108
109const CONNECTION_INIT_TIMEOUT: Duration = Duration::from_secs(5);
111
112const PING_INTERVAL: Duration = Duration::from_secs(30);
114
115#[derive(Clone)]
117pub struct SubscriptionState {
118 pub manager: Arc<SubscriptionManager>,
120 pub lifecycle: Arc<dyn SubscriptionLifecycle>,
122 pub max_subscriptions_per_connection: Option<u32>,
124 pub remote_subscription_fields: Arc<HashMap<String, String>>,
129 pub domain_registry: Option<Arc<DomainRegistry>>,
132 pub strict_tenant_validation: bool,
136 pub authorizer: Option<Arc<dyn Authorizer>>,
141 pub tenant_status_source: Option<Arc<dyn TenantStatusSource>>,
146 pub subscription_policies: Arc<HashMap<String, SubscriptionPolicy>>,
153 #[cfg(feature = "auth")]
158 pub identity_resolver: Option<Arc<crate::identity::IdentityResolver>>,
159 pub service_account_authenticator:
163 Option<Arc<crate::service_account::ServiceAccountAuthenticator>>,
164}
165
166impl SubscriptionState {
167 pub fn new(manager: Arc<SubscriptionManager>) -> Self {
169 Self {
170 manager,
171 lifecycle: Arc::new(crate::subscriptions::lifecycle::NoopLifecycle),
172 max_subscriptions_per_connection: None,
173 remote_subscription_fields: Arc::new(HashMap::new()),
174 domain_registry: None,
175 strict_tenant_validation: false,
176 authorizer: None,
177 tenant_status_source: None,
178 subscription_policies: Arc::new(HashMap::new()),
179 #[cfg(feature = "auth")]
180 identity_resolver: None,
181 service_account_authenticator: None,
182 }
183 }
184
185 #[must_use]
188 pub fn with_subscription_policies(
189 mut self,
190 policies: Arc<HashMap<String, SubscriptionPolicy>>,
191 ) -> Self {
192 self.subscription_policies = policies;
193 self
194 }
195
196 #[must_use]
199 pub fn with_service_account_authenticator(
200 mut self,
201 authenticator: Option<Arc<crate::service_account::ServiceAccountAuthenticator>>,
202 ) -> Self {
203 self.service_account_authenticator = authenticator;
204 self
205 }
206
207 #[cfg(feature = "auth")]
211 #[must_use]
212 pub fn with_identity_resolver(
213 mut self,
214 resolver: Option<Arc<crate::identity::IdentityResolver>>,
215 ) -> Self {
216 self.identity_resolver = resolver;
217 self
218 }
219
220 #[must_use]
224 pub fn with_tenant_status_source(
225 mut self,
226 source: Option<Arc<dyn TenantStatusSource>>,
227 ) -> Self {
228 self.tenant_status_source = source;
229 self
230 }
231
232 #[must_use]
237 pub fn with_authorizer(mut self, authorizer: Option<Arc<dyn Authorizer>>) -> Self {
238 self.authorizer = authorizer;
239 self
240 }
241
242 #[must_use]
249 pub fn with_tenant_context(
250 mut self,
251 domain_registry: Arc<DomainRegistry>,
252 strict_tenant_validation: bool,
253 ) -> Self {
254 self.domain_registry = Some(domain_registry);
255 self.strict_tenant_validation = strict_tenant_validation;
256 self
257 }
258
259 #[must_use]
261 pub fn with_lifecycle(mut self, lifecycle: Arc<dyn SubscriptionLifecycle>) -> Self {
262 self.lifecycle = lifecycle;
263 self
264 }
265
266 #[must_use]
268 pub const fn with_max_subscriptions(mut self, max: Option<u32>) -> Self {
269 self.max_subscriptions_per_connection = max;
270 self
271 }
272
273 #[must_use]
277 pub fn with_remote_subscription_fields(mut self, fields: HashMap<String, String>) -> Self {
278 self.remote_subscription_fields = Arc::new(fields);
279 self
280 }
281}
282
283#[must_use]
290pub fn build_subscription_policies(schema: &CompiledSchema) -> HashMap<String, SubscriptionPolicy> {
291 let mut policies = HashMap::new();
292 for sub in &schema.subscriptions {
293 if let Some(type_def) = schema.types.iter().find(|t| t.name.as_str() == sub.return_type) {
294 if let Some(policy) = &type_def.subscription_policy {
295 policies.insert(sub.name.clone(), policy.clone());
296 }
297 }
298 }
299 policies
300}
301
302async fn resolve_subscription_rls(
310 state: &SubscriptionState,
311 subscription_name: &str,
312 principal: Option<&SecurityContext>,
313) -> Result<Vec<(String, serde_json::Value)>, String> {
314 let Some(policy) = state.subscription_policies.get(subscription_name) else {
315 return Ok(Vec::new());
316 };
317
318 #[cfg(feature = "auth")]
324 let enriched = enrich_principal(state, principal).await;
325 #[cfg(feature = "auth")]
326 let effective = enriched.as_ref().or(principal);
327 #[cfg(not(feature = "auth"))]
328 let effective = principal;
329
330 derive_policy_conditions(policy, effective)
331}
332
333#[cfg(feature = "auth")]
338async fn enrich_principal(
339 state: &SubscriptionState,
340 principal: Option<&SecurityContext>,
341) -> Option<SecurityContext> {
342 let resolver = state.identity_resolver.as_ref()?;
343 let mut ctx = principal?.clone();
344 let _ = crate::identity::enrich_security_context(resolver, &mut ctx).await;
345 Some(ctx)
346}
347
348fn derive_policy_conditions(
353 policy: &SubscriptionPolicy,
354 principal: Option<&SecurityContext>,
355) -> Result<Vec<(String, serde_json::Value)>, String> {
356 let empty = HashMap::new();
357 let attributes = principal.map_or(&empty, |ctx| &ctx.attributes);
358 let roles: &[String] = principal.map_or(&[], |ctx| ctx.roles.as_slice());
359 match policy.derive(attributes, roles) {
360 OwnerCondition::Bypass => Ok(Vec::new()),
361 OwnerCondition::Eq { field, value } => Ok(vec![(field, value)]),
362 OwnerCondition::Refuse(reason) => Err(reason),
363 }
364}
365
366pub async fn subscription_handler(
373 headers: HeaderMap,
374 OptionalSecurityContext(security_context): OptionalSecurityContext,
375 ws: WebSocketUpgrade,
376 State(state): State<SubscriptionState>,
377) -> impl IntoResponse {
378 let protocol_header = headers.get("sec-websocket-protocol").and_then(|v| v.to_str().ok());
379
380 let protocol = match protocol_header {
381 None => WsProtocol::GraphqlTransportWs,
382 Some(header) => {
383 if let Some(p) = WsProtocol::from_header(Some(header)) {
384 p
385 } else {
386 warn!(header = %header, "Unknown WebSocket sub-protocol requested");
387 return axum::http::StatusCode::BAD_REQUEST.into_response();
388 }
389 },
390 };
391
392 let mut security_context = security_context;
398 if let Some(sa_auth) = state.service_account_authenticator.as_ref() {
399 match sa_auth.resolve(&headers, security_context.is_some()) {
400 crate::service_account::SaAuth::NoSecret => {},
401 crate::service_account::SaAuth::Authenticated(ctx) => security_context = Some(*ctx),
402 crate::service_account::SaAuth::Ambiguous
403 | crate::service_account::SaAuth::Unmatched => {
404 warn!(
405 "Subscription upgrade rejected: ambiguous or unmatched service-account secret"
406 );
407 return axum::http::StatusCode::UNAUTHORIZED.into_response();
408 },
409 }
410 }
411
412 let tenant_id = match resolve_subscription_tenant(security_context.as_ref(), &headers, &state) {
419 Ok(tenant_id) => tenant_id,
420 Err(e) => {
421 warn!(error = %e, "Subscription tenant resolution rejected the upgrade");
422 return axum::http::StatusCode::BAD_REQUEST.into_response();
423 },
424 };
425
426 ws.protocols([protocol.as_str()])
427 .on_upgrade(move |socket| {
428 handle_subscription_connection(socket, state, protocol, tenant_id, security_context)
429 })
430 .into_response()
431}
432
433fn resolve_subscription_tenant(
443 security_context: Option<&SecurityContext>,
444 headers: &HeaderMap,
445 state: &SubscriptionState,
446) -> fraiseql_error::Result<Option<String>> {
447 super::graphql::TenantKeyResolver::resolve(
448 security_context,
449 headers,
450 state.domain_registry.as_deref(),
451 state.strict_tenant_validation,
452 )
453}
454
455#[allow(clippy::cognitive_complexity)] async fn handle_subscription_connection(
458 socket: WebSocket,
459 state: SubscriptionState,
460 protocol: WsProtocol,
461 tenant_id: Option<String>,
462 principal: Option<SecurityContext>,
463) {
464 let connection_id = uuid::Uuid::new_v4().to_string();
465 let codec = ProtocolCodec::new(protocol);
466 info!(
467 connection_id = %connection_id,
468 protocol = %protocol.as_str(),
469 "WebSocket connection established"
470 );
471
472 let (mut sender, mut receiver) = socket.split();
473
474 let init_result = tokio::time::timeout(CONNECTION_INIT_TIMEOUT, async {
476 while let Some(msg) = receiver.next().await {
477 match msg {
478 Ok(Message::Text(text)) => {
479 if let Ok(client_msg) = codec.decode(&text) {
480 if client_msg.parsed_type() == Some(ClientMessageType::ConnectionInit) {
481 return Some(client_msg);
482 }
483 }
484 },
485 Ok(Message::Close(_)) => return None,
486 Err(e) => {
487 error!(error = %e, "WebSocket error during init");
488 return None;
489 },
490 _ => {},
491 }
492 }
493 None
494 })
495 .await;
496
497 let _init_payload = match init_result {
499 Ok(Some(msg)) => {
500 let params = msg.payload.clone().unwrap_or(serde_json::json!({}));
502 if let Err(reason) = state.lifecycle.on_connect(¶ms, &connection_id).await {
503 warn!(
504 connection_id = %connection_id,
505 reason = %reason,
506 "Lifecycle on_connect rejected connection"
507 );
508 WS_CONNECTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
509 let _ = sender
511 .send(Message::Close(Some(axum::extract::ws::CloseFrame {
512 code: 4400,
513 reason: reason.into(),
514 })))
515 .await;
516 return;
517 }
518
519 let ack = ServerMessage::connection_ack(None);
521 if let Err(send_err) = send_server_message(&codec, &mut sender, ack).await {
522 error!(connection_id = %connection_id, error = %send_err, "Failed to send connection_ack");
523 return;
524 }
525 WS_CONNECTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
526 info!(connection_id = %connection_id, "Connection initialized");
527 msg.payload
528 },
529 Ok(None) => {
530 warn!(connection_id = %connection_id, "Connection closed during init");
531 return;
532 },
533 Err(_) => {
534 warn!(connection_id = %connection_id, "Connection init timeout");
535 let _ = sender
537 .send(Message::Close(Some(axum::extract::ws::CloseFrame {
538 code: CloseCode::ConnectionInitTimeout.code(),
539 reason: CloseCode::ConnectionInitTimeout.reason().into(),
540 })))
541 .await;
542 return;
543 },
544 };
545
546 let mut active_operations: HashMap<String, SubscriptionId> = HashMap::new();
548
549 let (remote_msg_tx, mut remote_msg_rx) = tokio::sync::mpsc::channel::<ServerMessage>(64);
555
556 let mut event_receiver = state.manager.receiver();
558
559 let mut ping_interval = tokio::time::interval(PING_INTERVAL);
561 ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
562
563 loop {
583 tokio::select! {
584 msg = receiver.next() => {
585 match msg {
586 Some(Ok(Message::Text(text))) => {
587 if let Err(close_code) = handle_client_message(
588 &text,
589 &connection_id,
590 &state,
591 &codec,
592 &mut active_operations,
593 remote_msg_tx.clone(),
594 &mut sender,
595 tenant_id.as_deref(),
596 principal.as_ref(),
597 ).await {
598 let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
600 code: close_code.code(),
601 reason: close_code.reason().into(),
602 }))).await;
603 break;
604 }
605 }
606 Some(Ok(Message::Ping(data))) => {
607 let _ = sender.send(Message::Pong(data)).await;
609 }
610 Some(Ok(Message::Close(_))) => {
611 info!(connection_id = %connection_id, "Client closed connection");
612 break;
613 }
614 Some(Err(e)) => {
615 error!(connection_id = %connection_id, error = %e, "WebSocket error");
616 break;
617 }
618 None => {
619 info!(connection_id = %connection_id, "WebSocket stream ended");
620 break;
621 }
622 _ => {}
623 }
624 }
625
626 event = event_receiver.recv() => {
627 match event {
628 Ok(payload) => {
629 let tenant_matches = match (
635 tenant_id.as_deref(),
636 payload.event.tenant_id.as_deref(),
637 ) {
638 (Some(conn_tid), Some(evt_tid)) => conn_tid == evt_tid,
639 _ => true, };
641 let tenant_active = match (
645 tenant_id.as_deref(),
646 state.tenant_status_source.as_ref(),
647 ) {
648 (Some(tid), Some(src)) => !src.is_suspended(tid),
649 _ => true,
650 };
651 if tenant_matches && tenant_active {
652 if let Some((op_id, _)) = active_operations
653 .iter()
654 .find(|(_, sub_id)| **sub_id == payload.subscription_id)
655 {
656 let msg = create_next_message(op_id, &payload);
657 if send_server_message(&codec, &mut sender, msg).await.is_err() {
658 warn!(connection_id = %connection_id, "Failed to send event");
659 break;
660 }
661 }
662 }
663 }
664 Err(broadcast::error::RecvError::Lagged(n)) => {
665 warn!(connection_id = %connection_id, lagged = n, "Event receiver lagged");
666 }
667 Err(broadcast::error::RecvError::Closed) => {
668 error!(connection_id = %connection_id, "Event channel closed");
669 break;
670 }
671 }
672 }
673
674 remote_msg = remote_msg_rx.recv() => {
675 if let Some(msg) = remote_msg {
679 if send_server_message(&codec, &mut sender, msg).await.is_err() {
680 warn!(connection_id = %connection_id, "Failed to send remote subscription message");
681 break;
682 }
683 }
684 }
685
686 _ = ping_interval.tick() => {
687 let msg = ServerMessage::ping(None);
688 if send_server_message(&codec, &mut sender, msg).await.is_err() {
689 warn!(connection_id = %connection_id, "Failed to send ping/keepalive");
690 break;
691 }
692 }
693 }
694 }
695
696 state.manager.unsubscribe_connection(&connection_id);
698 state.lifecycle.on_disconnect(&connection_id).await;
699 info!(connection_id = %connection_id, "WebSocket connection closed");
700}
701
702#[allow(clippy::cognitive_complexity)] #[allow(clippy::too_many_arguments)] async fn handle_client_message(
711 text: &str,
712 connection_id: &str,
713 state: &SubscriptionState,
714 codec: &ProtocolCodec,
715 active_operations: &mut HashMap<String, SubscriptionId>,
716 remote_msg_tx: tokio::sync::mpsc::Sender<ServerMessage>,
717 sender: &mut futures::stream::SplitSink<WebSocket, Message>,
718 tenant_id: Option<&str>,
719 principal: Option<&SecurityContext>,
720) -> Result<(), CloseCode> {
721 #[cfg(not(feature = "federation"))]
724 let _ = &remote_msg_tx;
725
726 let client_msg: ClientMessage = codec.decode(text).map_err(|e| {
727 warn!(error = %e, "Failed to parse client message");
728 CloseCode::ProtocolError
729 })?;
730
731 match client_msg.parsed_type() {
732 Some(ClientMessageType::Ping) => {
733 let pong = ServerMessage::pong(client_msg.payload);
734 let _ = send_server_message(codec, sender, pong).await;
736 },
737
738 Some(ClientMessageType::Pong) => {
739 debug!(connection_id = %connection_id, "Received pong");
740 },
741
742 Some(ClientMessageType::Subscribe) => {
743 let payload: SubscribePayload = client_msg.subscription_payload().ok_or_else(|| {
744 warn!("Invalid subscribe payload");
745 CloseCode::ProtocolError
746 })?;
747
748 let op_id = client_msg.id.ok_or_else(|| {
749 warn!("Subscribe message missing operation ID");
750 CloseCode::ProtocolError
751 })?;
752
753 if active_operations.contains_key(&op_id) {
755 warn!(operation_id = %op_id, "Duplicate operation ID");
756 return Err(CloseCode::SubscriberAlreadyExists);
757 }
758
759 if let Some(max) = state.max_subscriptions_per_connection {
761 if active_operations.len() >= max as usize {
762 warn!(
763 connection_id = %connection_id,
764 active = active_operations.len(),
765 max = max,
766 "Subscription limit reached"
767 );
768 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
769 let error = ServerMessage::error(
770 &op_id,
771 vec![GraphQLError::with_code(
772 format!("Maximum subscriptions per connection ({max}) reached"),
773 "SUBSCRIPTION_LIMIT_REACHED",
774 )],
775 );
776 if let Err(e) = send_server_message(codec, sender, error).await {
777 debug!(connection_id = %connection_id, error = %e, "Could not send subscription limit error to client");
778 }
779 return Ok(());
780 }
781 }
782
783 let Some(subscription_name) = extract_subscription_name(&payload.query) else {
785 let error = ServerMessage::error(
786 &op_id,
787 vec![GraphQLError::with_code(
788 "Could not parse subscription query",
789 "PARSE_ERROR",
790 )],
791 );
792 if let Err(e) = send_server_message(codec, sender, error).await {
793 debug!(connection_id = %connection_id, error = %e, "Could not send parse error to client");
794 }
795 return Ok(());
796 };
797
798 let variables_value = serde_json::to_value(&payload.variables)
801 .expect("HashMap<String, serde_json::Value> serialization is infallible");
802
803 if let Some(authorizer) = state.authorizer.as_ref() {
809 let ops = [(OperationKind::Subscription, subscription_name.clone())];
810 if let Err(err) =
811 enforce_authz(authorizer.as_ref(), principal, &ops, Some(&variables_value))
812 {
813 warn!(
814 connection_id = %connection_id,
815 subscription = %subscription_name,
816 "Operation authorizer denied the subscription"
817 );
818 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
819 let error = ServerMessage::error(
820 &op_id,
821 vec![GraphQLError::with_code(err.to_string(), "FORBIDDEN")],
822 );
823 if let Err(e) = send_server_message(codec, sender, error).await {
824 debug!(connection_id = %connection_id, error = %e, "Could not send authorization denial to client");
825 }
826 return Ok(());
827 }
828 }
829
830 if let Err(reason) = state
831 .lifecycle
832 .on_subscribe(&subscription_name, &variables_value, connection_id)
833 .await
834 {
835 warn!(
836 connection_id = %connection_id,
837 subscription = %subscription_name,
838 reason = %reason,
839 "Lifecycle on_subscribe rejected subscription"
840 );
841 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
842 let error = ServerMessage::error(
843 &op_id,
844 vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REJECTED")],
845 );
846 if let Err(e) = send_server_message(codec, sender, error).await {
847 debug!(connection_id = %connection_id, error = %e, "Could not send subscription rejection to client");
848 }
849 return Ok(());
850 }
851
852 #[cfg(feature = "federation")]
854 if let Some(subgraph_url) = state.remote_subscription_fields.get(&subscription_name) {
855 use fraiseql_federation::subscription_forwarder::{
856 ForwardedEvent, SubscriptionForwarder,
857 };
858
859 match SubscriptionForwarder::new(subgraph_url) {
860 Ok(forwarder) => {
861 let (event_tx, mut event_rx) =
863 tokio::sync::mpsc::channel::<ForwardedEvent>(32);
864
865 let fwd_op = op_id.clone();
867 let fwd_query = payload.query.clone();
868 let fwd_vars = variables_value.clone();
869 tokio::spawn(async move {
870 if let Err(e) =
871 forwarder.forward(&fwd_op, &fwd_query, fwd_vars, event_tx).await
872 {
873 warn!(error = %e, "Remote subscription forwarder failed");
874 }
875 });
876
877 let relay_op = op_id.clone();
879 let relay_tx = remote_msg_tx.clone();
880 tokio::spawn(async move {
881 while let Some(event) = event_rx.recv().await {
882 let server_msg = match event {
883 ForwardedEvent::Next(data) => {
884 ServerMessage::next(&relay_op, data)
885 },
886 ForwardedEvent::Error(errors) => {
887 let errors_vec = errors.as_array().map_or_else(
888 || {
889 vec![GraphQLError::with_code(
890 errors.to_string(),
891 "REMOTE_ERROR",
892 )]
893 },
894 |arr| {
895 arr.iter()
896 .map(|e| {
897 GraphQLError::with_code(
898 e.get("message")
899 .and_then(|v| v.as_str())
900 .unwrap_or("Remote subgraph error"),
901 "REMOTE_ERROR",
902 )
903 })
904 .collect()
905 },
906 );
907 ServerMessage::error(&relay_op, errors_vec)
908 },
909 ForwardedEvent::Complete => ServerMessage::complete(&relay_op),
910 };
911 if relay_tx.send(server_msg).await.is_err() {
912 break; }
914 }
915 });
916
917 WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
918 info!(
919 connection_id = %connection_id,
920 operation_id = %op_id,
921 subscription = %subscription_name,
922 "Subscription forwarded to remote subgraph"
923 );
924 return Ok(());
925 },
926 Err(e) => {
927 let error = ServerMessage::error(
928 &op_id,
929 vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
930 );
931 if let Err(send_err) = send_server_message(codec, sender, error).await {
932 debug!(connection_id = %connection_id, error = %send_err, "Could not send forwarding error to client");
933 }
934 return Ok(());
935 },
936 }
937 }
938
939 if let Some(server_tid) = tenant_id {
941 if let Some(client_tid) = variables_value.get("tenant_id").and_then(|v| v.as_str())
942 {
943 if client_tid != server_tid {
944 let error = ServerMessage::error(
945 &op_id,
946 vec![GraphQLError::with_code(
947 format!(
948 "Tenant mismatch: client provided '{client_tid}', server resolved '{server_tid}'"
949 ),
950 "TENANT_MISMATCH",
951 )],
952 );
953 if let Err(send_err) = send_server_message(codec, sender, error).await {
954 debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant mismatch error to client");
955 }
956 return Ok(());
957 }
958 }
959 }
960
961 if let (Some(tid), Some(src)) = (tenant_id, state.tenant_status_source.as_ref()) {
964 if src.is_suspended(tid) {
965 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
966 let error = ServerMessage::error(
967 &op_id,
968 vec![GraphQLError::with_code(
969 format!("Tenant '{tid}' is suspended"),
970 "TENANT_SUSPENDED",
971 )],
972 );
973 if let Err(send_err) = send_server_message(codec, sender, error).await {
974 debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant-suspended error to client");
975 }
976 return Ok(());
977 }
978 }
979
980 let rls_conditions = match resolve_subscription_rls(
985 state,
986 &subscription_name,
987 principal,
988 )
989 .await
990 {
991 Ok(conditions) => conditions,
992 Err(reason) => {
993 WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
994 warn!(
995 connection_id = %connection_id,
996 subscription = %subscription_name,
997 "Row-visibility policy refused the subscription (fail-closed)"
998 );
999 let error = ServerMessage::error(
1000 &op_id,
1001 vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REFUSED")],
1002 );
1003 if let Err(send_err) = send_server_message(codec, sender, error).await {
1004 debug!(connection_id = %connection_id, error = %send_err, "Could not send row-visibility refusal to client");
1005 }
1006 return Ok(());
1007 },
1008 };
1009
1010 let mut context = serde_json::json!({});
1012 if let Some(tid) = tenant_id {
1013 context["tenant_id"] = serde_json::Value::String(tid.to_string());
1014 }
1015
1016 match state.manager.subscribe_with_rls(
1020 &subscription_name,
1021 context,
1022 variables_value,
1023 connection_id,
1024 rls_conditions,
1025 ) {
1026 Ok(sub_id) => {
1027 active_operations.insert(op_id.clone(), sub_id);
1028 WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
1029 info!(
1030 connection_id = %connection_id,
1031 operation_id = %op_id,
1032 subscription = %subscription_name,
1033 "Subscription started"
1034 );
1035 },
1036 Err(e) => {
1037 let error = ServerMessage::error(
1038 &op_id,
1039 vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
1040 );
1041 if let Err(send_err) = send_server_message(codec, sender, error).await {
1042 debug!(connection_id = %connection_id, error = %send_err, "Could not send subscription error to client");
1043 }
1044 },
1045 }
1046 },
1047
1048 Some(ClientMessageType::Complete) => {
1049 let op_id = client_msg.id.ok_or_else(|| {
1050 warn!("Complete message missing operation ID");
1051 CloseCode::ProtocolError
1052 })?;
1053
1054 if let Some(sub_id) = active_operations.remove(&op_id) {
1055 if let Err(e) = state.manager.unsubscribe(sub_id) {
1056 warn!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe; subscription may be leaked");
1057 }
1058 state.lifecycle.on_unsubscribe(&op_id, connection_id).await;
1059 info!(
1060 connection_id = %connection_id,
1061 operation_id = %op_id,
1062 "Subscription completed"
1063 );
1064 }
1065 },
1066
1067 Some(ClientMessageType::ConnectionInit) => {
1068 warn!(connection_id = %connection_id, "Duplicate connection_init");
1069 return Err(CloseCode::TooManyInitRequests);
1070 },
1071
1072 None => {
1073 warn!(message_type = %client_msg.message_type, "Unknown message type");
1074 },
1075 _ => {
1077 warn!(message_type = %client_msg.message_type, "Unrecognized message type");
1078 },
1079 }
1080
1081 Ok(())
1082}
1083
1084async fn send_server_message(
1086 codec: &ProtocolCodec,
1087 sender: &mut futures::stream::SplitSink<WebSocket, Message>,
1088 msg: ServerMessage,
1089) -> Result<(), String> {
1090 match codec.encode(&msg) {
1091 Ok(Some(json)) => sender.send(Message::Text(json.into())).await.map_err(|e| e.to_string()),
1092 Ok(None) => Ok(()), Err(e) => Err(e.to_string()),
1094 }
1095}
1096
1097fn create_next_message(operation_id: &str, payload: &SubscriptionPayload) -> ServerMessage {
1104 let data = serde_json::json!({
1105 payload.subscription_name.clone(): payload.data
1106 });
1107 match &payload.event.change_spine {
1108 Some(envelope) => {
1109 let extensions = serde_json::json!({ "changeSpine": envelope });
1110 ServerMessage::next_with_extensions(operation_id, data, extensions)
1111 },
1112 None => ServerMessage::next(operation_id, data),
1113 }
1114}
1115
1116pub(crate) fn extract_subscription_name(query: &str) -> Option<String> {
1118 let query = query.trim();
1119
1120 let sub_idx = query.find("subscription")?;
1121 let after_sub = &query[sub_idx + "subscription".len()..];
1122
1123 let brace_idx = after_sub.find('{')?;
1124 let after_brace = after_sub[brace_idx + 1..].trim_start();
1125
1126 let name_end = after_brace
1127 .find(|c: char| !c.is_alphanumeric() && c != '_')
1128 .unwrap_or(after_brace.len());
1129
1130 if name_end == 0 {
1131 return None;
1132 }
1133
1134 Some(after_brace[..name_end].to_string())
1135}
1136
1137#[cfg(test)]
1138mod tests;