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::{
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
66// ── Subscription metrics (module-level atomics) ──────────────────────
67
68static 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/// Subscription metrics for Prometheus export.
74#[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/// Reset all subscription counters to zero.
85///
86/// Call this at the start of each test that checks counter values to avoid
87/// cross-test interference from the module-level statics.
88#[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
96/// Snapshot of subscription counters.
97pub struct SubscriptionMetrics {
98    /// Total `WebSocket` connections accepted (after `on_connect`).
99    pub connections_accepted:   u64,
100    /// Total `WebSocket` connections rejected by lifecycle hook.
101    pub connections_rejected:   u64,
102    /// Total subscriptions accepted (after `on_subscribe`).
103    pub subscriptions_accepted: u64,
104    /// Total subscriptions rejected (by hook or limit).
105    pub subscriptions_rejected: u64,
106}
107
108/// Connection initialization timeout (5 seconds per graphql-ws spec).
109const CONNECTION_INIT_TIMEOUT: Duration = Duration::from_secs(5);
110
111/// Ping/keepalive interval.
112const PING_INTERVAL: Duration = Duration::from_secs(30);
113
114/// State for subscription `WebSocket` handler.
115#[derive(Clone)]
116pub struct SubscriptionState {
117    /// Subscription manager.
118    pub manager: Arc<SubscriptionManager>,
119    /// Lifecycle hooks.
120    pub lifecycle: Arc<dyn SubscriptionLifecycle>,
121    /// Maximum subscriptions per connection (`None` = unlimited).
122    pub max_subscriptions_per_connection: Option<u32>,
123    /// Subscription fields owned by remote subgraphs.
124    ///
125    /// Maps root subscription field name to the subgraph `WebSocket` URL.
126    /// Empty when federation is disabled or no remote subscription fields are declared.
127    pub remote_subscription_fields: Arc<HashMap<String, String>>,
128    /// Host-header → tenant-key domain registry. `None` until a host binary
129    /// installs one (mirrors `AppState::domain_registry`).
130    pub domain_registry: Option<Arc<DomainRegistry>>,
131    /// Reject conflicting tenant sources (JWT vs `X-Tenant-ID` vs Host) on the
132    /// upgrade. Driven by `schema.has_rls_configured()`, mirroring the GraphQL
133    /// handler's strict tenant validation.
134    pub strict_tenant_validation: bool,
135    /// Optional operation-level authorizer (#422). When set, each subscription is
136    /// authorized at establishment with [`OperationKind::Subscription`], the
137    /// subscription field name, and the connection's principal. `None` until a host
138    /// binary installs one (from `Executor::config().authorizer`).
139    pub authorizer: Option<Arc<dyn Authorizer>>,
140    /// Optional tenant-status source (M-tenant-ws-suspended). When set, a new
141    /// subscription whose resolved tenant is suspended is rejected, and event
142    /// delivery to a connection whose tenant is suspended is paused. `None` until
143    /// a host binary installs a multi-tenant registry.
144    pub tenant_status_source: Option<Arc<dyn TenantStatusSource>>,
145}
146
147impl SubscriptionState {
148    /// Create new subscription state.
149    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    /// Install the tenant-status source (M-tenant-ws-suspended). When set, new
163    /// subscriptions for a suspended tenant are rejected and event delivery to a
164    /// suspended tenant is paused. Typically the `TenantExecutorRegistry`.
165    #[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    /// Install the operation-level authorizer (#422). When set, every subscription is
175    /// authorized at establishment; a `Deny` (or any policy error) rejects the
176    /// subscription with a `FORBIDDEN` GraphQL-WS error. Typically populated from
177    /// `Executor::config().authorizer`.
178    #[must_use]
179    pub fn with_authorizer(mut self, authorizer: Option<Arc<dyn Authorizer>>) -> Self {
180        self.authorizer = authorizer;
181        self
182    }
183
184    /// Install the tenant-resolution context — the Host-header domain registry
185    /// and strict-validation flag — so the subscription upgrade dispatches the
186    /// tenant key the same way the GraphQL handler does (JWT `tenant_id` >
187    /// `X-Tenant-ID` header > Host domain, with cross-source conflict rejection
188    /// when `strict_tenant_validation` is set). See
189    /// [`crate::routes::graphql::TenantKeyResolver`].
190    #[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    /// Set lifecycle hooks.
202    #[must_use]
203    pub fn with_lifecycle(mut self, lifecycle: Arc<dyn SubscriptionLifecycle>) -> Self {
204        self.lifecycle = lifecycle;
205        self
206    }
207
208    /// Set maximum subscriptions per connection.
209    #[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    /// Set remote subscription fields (federation passthrough).
216    ///
217    /// Maps subscription field names to the owning subgraph's `WebSocket` URL.
218    #[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
225/// `WebSocket` upgrade handler for subscriptions.
226///
227/// Negotiates the `WebSocket` sub-protocol from the `Sec-WebSocket-Protocol`
228/// header. Supports `graphql-transport-ws` (modern) and `graphql-ws` (legacy).
229/// Defaults to `graphql-transport-ws` when no header is present.
230/// Returns `400 Bad Request` for unrecognised protocols.
231pub 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    // Resolve the tenant key exactly as the GraphQL handler does: the JWT
252    // `tenant_id` (trusted) takes precedence over the `X-Tenant-ID` header and
253    // the Host-domain lookup, and conflicting sources are rejected when strict
254    // validation is enabled. Previously this passed `None, None, false`,
255    // silently dropping JWT precedence, ignoring an installed domain registry,
256    // and disabling strict cross-source validation (#331).
257    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
272/// Resolve the subscription's tenant key, mirroring the GraphQL handler's
273/// dispatch (`routes/graphql/handler.rs`): JWT `tenant_id` > `X-Tenant-ID`
274/// header > Host-domain registry, with cross-source conflict rejection when the
275/// schema has RLS configured (`strict_tenant_validation`).
276///
277/// # Errors
278///
279/// Returns `FraiseQLError::Validation` when the `X-Tenant-ID` header is invalid,
280/// or — under `strict_tenant_validation` — when the resolved sources disagree.
281fn 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/// Handle a `WebSocket` subscription connection.
295#[allow(clippy::cognitive_complexity)] // Reason: WebSocket protocol state machine with message routing and lifecycle management
296async 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    // Wait for connection_init with timeout
314    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    // Handle init timeout or failure
337    let _init_payload = match init_result {
338        Ok(Some(msg)) => {
339            // Call lifecycle on_connect hook
340            let params = msg.payload.clone().unwrap_or(serde_json::json!({}));
341            if let Err(reason) = state.lifecycle.on_connect(&params, &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                // Best-effort: connection is already being terminated.
349                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            // Send connection_ack
359            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            // Best-effort: connection is already being terminated.
375            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    // Track active operations (operation_id -> subscription_id)
386    let mut active_operations: HashMap<String, SubscriptionId> = HashMap::new();
387
388    // Remote subscription message output channel.
389    //
390    // Forwarder tasks (federation feature) send pre-encoded ServerMessage values here.
391    // The channel is always present so the select! loop has a uniform branch regardless
392    // of whether any remote subscriptions are active.
393    let (remote_msg_tx, mut remote_msg_rx) = tokio::sync::mpsc::channel::<ServerMessage>(64);
394
395    // Subscribe to event broadcast
396    let mut event_receiver = state.manager.receiver();
397
398    // Ping/keepalive timer
399    let mut ping_interval = tokio::time::interval(PING_INTERVAL);
400    ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
401
402    // A44 — Token expiry re-check on long-lived subscriptions.
403    //
404    // JWTs validated at ConnectionInit may expire while the WebSocket is open.
405    // The check below should be added when the auth layer surfaces expiry data:
406    //
407    //   1. At ConnectionInit, extract the `exp` claim from the JWT and store it: `let
408    //      token_expires_at: Option<std::time::Instant> = extract_exp(&init_payload);`
409    //
410    //   2. In the select! loop (before processing each client message or broadcast event), check
411    //      expiry: ```rust,ignore if token_expires_at.is_some_and(|exp| std::time::Instant::now()
412    //      >= exp) { warn!(connection_id = %connection_id, "Token expired; closing WebSocket"); let
413    //      _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame { code:
414    //      CloseCode::Unauthorized.code(), reason: "Token expired".into(), }))).await; break; } ```
415    //
416    // This requires the lifecycle `on_connect` hook or the JWT middleware to return
417    // the expiry time, which is not yet threaded through `SubscriptionState`.
418    // Tracked as A44 in the remediation plan.
419
420    // Main message loop
421    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                            // Best-effort: connection is already being closed.
438                            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                        // Best-effort: if the connection is already dead the pong will fail.
447                        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                        // Defense-in-depth tenant guard: when both the connection and the
469                        // event carry an explicit tenant_id they must agree. Primary
470                        // isolation is already guaranteed by subscription_id UUIDs, but
471                        // this check catches any future path that introduces deterministic
472                        // subscription IDs (which could collide across tenants).
473                        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, // either side absent → no conflict
479                        };
480                        // M-tenant-ws-suspended: pause delivery while this
481                        // connection's tenant is suspended (re-checked per event so a
482                        // suspension mid-stream stops further delivery).
483                        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                // Remote subscription event forwarded by a federation forwarder task.
515                // The channel is always present; it only carries messages when the
516                // federation feature is enabled and remote subscriptions are active.
517                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    // Cleanup
536    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/// Handle a client message.
542///
543/// Returns `Ok(())` on success, or `Err(CloseCode)` if the connection should be closed.
544///
545/// `remote_msg_tx` is used by federation forwarder tasks to send pre-encoded
546/// `ServerMessage` values back to the client connection loop.
547#[allow(clippy::cognitive_complexity)] // Reason: WebSocket message dispatch with subscribe/unsubscribe/query protocol handling
548#[allow(clippy::too_many_arguments)] // Reason: WebSocket handler needs connection state, protocol codec, and tenant context
549async 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    // remote_msg_tx is only consumed inside the #[cfg(feature = "federation")] block below.
561    // When the feature is disabled the parameter goes unused; suppress the warning.
562    #[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            // Best-effort: if the connection is already dead the pong will fail.
574            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            // Check for duplicate operation ID
593            if active_operations.contains_key(&op_id) {
594                warn!(operation_id = %op_id, "Duplicate operation ID");
595                return Err(CloseCode::SubscriberAlreadyExists);
596            }
597
598            // Enforce per-connection subscription limit
599            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            // Extract subscription name from query
623            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            // Call lifecycle on_subscribe hook
638            // HashMap<String, Value> serialization is infallible; the error path cannot occur.
639            let variables_value = serde_json::to_value(&payload.variables)
640                .expect("HashMap<String, serde_json::Value> serialization is infallible");
641
642            // #422: operation-level authorization at subscription establishment.
643            // The per-event delivery does not route through the executor, so the
644            // subscription is authorized once, here, with the connection's principal
645            // (or `None` when anonymous). Fail-closed: a `Deny` or any policy error
646            // rejects the subscription with a `FORBIDDEN` GraphQL-WS error.
647            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            // Forward to remote subgraph when the subscription field is owned remotely.
692            #[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                        // Create a channel so the forwarder task can send us raw events.
701                        let (event_tx, mut event_rx) =
702                            tokio::sync::mpsc::channel::<ForwardedEvent>(32);
703
704                        // Task 1: run the WebSocket forwarder (sends ForwardedEvent to event_tx).
705                        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                        // Task 2: relay ForwardedEvent → ServerMessage → client.
717                        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; // Client disconnected
752                                }
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            // Validate client-provided tenant variable against server-resolved
779            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            // M-tenant-ws-suspended: refuse to start a subscription whose tenant
801            // is suspended, mirroring the GraphQL data plane's 503 response.
802            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            // Build context with server-resolved tenant_id
820            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            // Subscribe locally (field is owned by this subgraph)
826            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        // Reason: non_exhaustive requires catch-all for cross-crate matches
882        _ => {
883            warn!(message_type = %client_msg.message_type, "Unrecognized message type");
884        },
885    }
886
887    Ok(())
888}
889
890/// Send a server message through the codec, handling protocol translation.
891async 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(()), // Message suppressed by codec (e.g. pong in legacy mode)
899        Err(e) => Err(e.to_string()),
900    }
901}
902
903/// Create a "next" message for a subscription event.
904fn create_next_message(operation_id: &str, payload: &SubscriptionPayload) -> ServerMessage {
905    let data = serde_json::json!({
906        payload.subscription_name.clone(): payload.data
907    });
908    ServerMessage::next(operation_id, data)
909}
910
911/// Extract subscription name from a GraphQL subscription query.
912pub(crate) fn extract_subscription_name(query: &str) -> Option<String> {
913    let query = query.trim();
914
915    let sub_idx = query.find("subscription")?;
916    let after_sub = &query[sub_idx + "subscription".len()..];
917
918    let brace_idx = after_sub.find('{')?;
919    let after_brace = after_sub[brace_idx + 1..].trim_start();
920
921    let name_end = after_brace
922        .find(|c: char| !c.is_alphanumeric() && c != '_')
923        .unwrap_or(after_brace.len());
924
925    if name_end == 0 {
926        return None;
927    }
928
929    Some(after_brace[..name_end].to_string())
930}
931
932#[cfg(test)]
933mod tests;