Skip to main content

fraiseql_server/routes/
subscriptions.rs

1//! `WebSocket` subscription handler with protocol negotiation.
2//!
3//! Supports both the modern `graphql-transport-ws` protocol and the legacy
4//! `graphql-ws` (Apollo subscriptions-transport-ws) protocol. Protocol
5//! selection happens during the `WebSocket` upgrade via the `Sec-WebSocket-Protocol`
6//! header.
7//!
8//! # Lifecycle Hooks
9//!
10//! Configurable callbacks are invoked at key points in the subscription
11//! lifecycle: `on_connect`, `on_disconnect`, `on_subscribe`, `on_unsubscribe`.
12//!
13//! # Example
14//!
15//! ```text
16//! // Requires: running server with initialized subscription manager.
17//! use fraiseql_server::routes::subscriptions::{subscription_handler, SubscriptionState};
18//!
19//! let state = SubscriptionState::new(subscription_manager);
20//!
21//! let app = Router::new()
22//!     .route("/ws", get(subscription_handler))
23//!     .with_state(state);
24//! ```
25
26use 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
58// ── Subscription metrics (module-level atomics) ──────────────────────
59
60static 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/// Subscription metrics for Prometheus export.
66#[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/// Reset all subscription counters to zero.
77///
78/// Call this at the start of each test that checks counter values to avoid
79/// cross-test interference from the module-level statics.
80#[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
88/// Snapshot of subscription counters.
89pub struct SubscriptionMetrics {
90    /// Total `WebSocket` connections accepted (after `on_connect`).
91    pub connections_accepted:   u64,
92    /// Total `WebSocket` connections rejected by lifecycle hook.
93    pub connections_rejected:   u64,
94    /// Total subscriptions accepted (after `on_subscribe`).
95    pub subscriptions_accepted: u64,
96    /// Total subscriptions rejected (by hook or limit).
97    pub subscriptions_rejected: u64,
98}
99
100/// Connection initialization timeout (5 seconds per graphql-ws spec).
101const CONNECTION_INIT_TIMEOUT: Duration = Duration::from_secs(5);
102
103/// Ping/keepalive interval.
104const PING_INTERVAL: Duration = Duration::from_secs(30);
105
106/// State for subscription `WebSocket` handler.
107#[derive(Clone)]
108pub struct SubscriptionState {
109    /// Subscription manager.
110    pub manager: Arc<SubscriptionManager>,
111    /// Lifecycle hooks.
112    pub lifecycle: Arc<dyn SubscriptionLifecycle>,
113    /// Maximum subscriptions per connection (`None` = unlimited).
114    pub max_subscriptions_per_connection: Option<u32>,
115    /// Subscription fields owned by remote subgraphs.
116    ///
117    /// Maps root subscription field name to the subgraph `WebSocket` URL.
118    /// Empty when federation is disabled or no remote subscription fields are declared.
119    pub remote_subscription_fields: Arc<HashMap<String, String>>,
120}
121
122impl SubscriptionState {
123    /// Create new subscription state.
124    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    /// Set lifecycle hooks.
134    #[must_use]
135    pub fn with_lifecycle(mut self, lifecycle: Arc<dyn SubscriptionLifecycle>) -> Self {
136        self.lifecycle = lifecycle;
137        self
138    }
139
140    /// Set maximum subscriptions per connection.
141    #[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    /// Set remote subscription fields (federation passthrough).
148    ///
149    /// Maps subscription field names to the owning subgraph's `WebSocket` URL.
150    #[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
157/// `WebSocket` upgrade handler for subscriptions.
158///
159/// Negotiates the `WebSocket` sub-protocol from the `Sec-WebSocket-Protocol`
160/// header. Supports `graphql-transport-ws` (modern) and `graphql-ws` (legacy).
161/// Defaults to `graphql-transport-ws` when no header is present.
162/// Returns `400 Bad Request` for unrecognised protocols.
163pub 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    // Resolve tenant from headers (same as GraphQL handler)
183    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/// Handle a `WebSocket` subscription connection.
195#[allow(clippy::cognitive_complexity)] // Reason: WebSocket protocol state machine with message routing and lifecycle management
196async 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    // Wait for connection_init with timeout
213    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    // Handle init timeout or failure
236    let _init_payload = match init_result {
237        Ok(Some(msg)) => {
238            // Call lifecycle on_connect hook
239            let params = msg.payload.clone().unwrap_or(serde_json::json!({}));
240            if let Err(reason) = state.lifecycle.on_connect(&params, &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                // Best-effort: connection is already being terminated.
248                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            // Send connection_ack
258            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            // Best-effort: connection is already being terminated.
274            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    // Track active operations (operation_id -> subscription_id)
285    let mut active_operations: HashMap<String, SubscriptionId> = HashMap::new();
286
287    // Remote subscription message output channel.
288    //
289    // Forwarder tasks (federation feature) send pre-encoded ServerMessage values here.
290    // The channel is always present so the select! loop has a uniform branch regardless
291    // of whether any remote subscriptions are active.
292    let (remote_msg_tx, mut remote_msg_rx) = tokio::sync::mpsc::channel::<ServerMessage>(64);
293
294    // Subscribe to event broadcast
295    let mut event_receiver = state.manager.receiver();
296
297    // Ping/keepalive timer
298    let mut ping_interval = tokio::time::interval(PING_INTERVAL);
299    ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
300
301    // A44 — Token expiry re-check on long-lived subscriptions.
302    //
303    // JWTs validated at ConnectionInit may expire while the WebSocket is open.
304    // The check below should be added when the auth layer surfaces expiry data:
305    //
306    //   1. At ConnectionInit, extract the `exp` claim from the JWT and store it: `let
307    //      token_expires_at: Option<std::time::Instant> = extract_exp(&init_payload);`
308    //
309    //   2. In the select! loop (before processing each client message or broadcast event), check
310    //      expiry: ```rust,ignore if token_expires_at.is_some_and(|exp| std::time::Instant::now()
311    //      >= exp) { warn!(connection_id = %connection_id, "Token expired; closing WebSocket"); let
312    //      _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame { code:
313    //      CloseCode::Unauthorized.code(), reason: "Token expired".into(), }))).await; break; } ```
314    //
315    // This requires the lifecycle `on_connect` hook or the JWT middleware to return
316    // the expiry time, which is not yet threaded through `SubscriptionState`.
317    // Tracked as A44 in the remediation plan.
318
319    // Main message loop
320    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                            // Best-effort: connection is already being closed.
336                            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                        // Best-effort: if the connection is already dead the pong will fail.
345                        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                        // Defense-in-depth tenant guard: when both the connection and the
367                        // event carry an explicit tenant_id they must agree. Primary
368                        // isolation is already guaranteed by subscription_id UUIDs, but
369                        // this check catches any future path that introduces deterministic
370                        // subscription IDs (which could collide across tenants).
371                        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, // either side absent → no conflict
377                        };
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                // Remote subscription event forwarded by a federation forwarder task.
403                // The channel is always present; it only carries messages when the
404                // federation feature is enabled and remote subscriptions are active.
405                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    // Cleanup
424    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/// Handle a client message.
430///
431/// Returns `Ok(())` on success, or `Err(CloseCode)` if the connection should be closed.
432///
433/// `remote_msg_tx` is used by federation forwarder tasks to send pre-encoded
434/// `ServerMessage` values back to the client connection loop.
435#[allow(clippy::cognitive_complexity)] // Reason: WebSocket message dispatch with subscribe/unsubscribe/query protocol handling
436#[allow(clippy::too_many_arguments)] // Reason: WebSocket handler needs connection state, protocol codec, and tenant context
437async 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    // remote_msg_tx is only consumed inside the #[cfg(feature = "federation")] block below.
448    // When the feature is disabled the parameter goes unused; suppress the warning.
449    #[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            // Best-effort: if the connection is already dead the pong will fail.
461            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            // Check for duplicate operation ID
480            if active_operations.contains_key(&op_id) {
481                warn!(operation_id = %op_id, "Duplicate operation ID");
482                return Err(CloseCode::SubscriberAlreadyExists);
483            }
484
485            // Enforce per-connection subscription limit
486            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            // Extract subscription name from query
510            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            // Call lifecycle on_subscribe hook
525            // HashMap<String, Value> serialization is infallible; the error path cannot occur.
526            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            // Forward to remote subgraph when the subscription field is owned remotely.
551            #[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                        // Create a channel so the forwarder task can send us raw events.
560                        let (event_tx, mut event_rx) =
561                            tokio::sync::mpsc::channel::<ForwardedEvent>(32);
562
563                        // Task 1: run the WebSocket forwarder (sends ForwardedEvent to event_tx).
564                        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                        // Task 2: relay ForwardedEvent → ServerMessage → client.
576                        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; // Client disconnected
611                                }
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            // Validate client-provided tenant variable against server-resolved
638            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            // Build context with server-resolved tenant_id
660            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            // Subscribe locally (field is owned by this subgraph)
666            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        // Reason: non_exhaustive requires catch-all for cross-crate matches
722        _ => {
723            warn!(message_type = %client_msg.message_type, "Unrecognized message type");
724        },
725    }
726
727    Ok(())
728}
729
730/// Send a server message through the codec, handling protocol translation.
731async 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(()), // Message suppressed by codec (e.g. pong in legacy mode)
739        Err(e) => Err(e.to_string()),
740    }
741}
742
743/// Create a "next" message for a subscription event.
744fn 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
751/// Extract subscription name from a GraphQL subscription query.
752pub(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}