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,
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}
141
142impl SubscriptionState {
143    /// Create new subscription state.
144    pub fn new(manager: Arc<SubscriptionManager>) -> Self {
145        Self {
146            manager,
147            lifecycle: Arc::new(crate::subscriptions::lifecycle::NoopLifecycle),
148            max_subscriptions_per_connection: None,
149            remote_subscription_fields: Arc::new(HashMap::new()),
150            domain_registry: None,
151            strict_tenant_validation: false,
152            authorizer: None,
153        }
154    }
155
156    /// Install the operation-level authorizer (#422). When set, every subscription is
157    /// authorized at establishment; a `Deny` (or any policy error) rejects the
158    /// subscription with a `FORBIDDEN` GraphQL-WS error. Typically populated from
159    /// `Executor::config().authorizer`.
160    #[must_use]
161    pub fn with_authorizer(mut self, authorizer: Option<Arc<dyn Authorizer>>) -> Self {
162        self.authorizer = authorizer;
163        self
164    }
165
166    /// Install the tenant-resolution context — the Host-header domain registry
167    /// and strict-validation flag — so the subscription upgrade dispatches the
168    /// tenant key the same way the GraphQL handler does (JWT `tenant_id` >
169    /// `X-Tenant-ID` header > Host domain, with cross-source conflict rejection
170    /// when `strict_tenant_validation` is set). See
171    /// [`crate::routes::graphql::TenantKeyResolver`].
172    #[must_use]
173    pub fn with_tenant_context(
174        mut self,
175        domain_registry: Arc<DomainRegistry>,
176        strict_tenant_validation: bool,
177    ) -> Self {
178        self.domain_registry = Some(domain_registry);
179        self.strict_tenant_validation = strict_tenant_validation;
180        self
181    }
182
183    /// Set lifecycle hooks.
184    #[must_use]
185    pub fn with_lifecycle(mut self, lifecycle: Arc<dyn SubscriptionLifecycle>) -> Self {
186        self.lifecycle = lifecycle;
187        self
188    }
189
190    /// Set maximum subscriptions per connection.
191    #[must_use]
192    pub const fn with_max_subscriptions(mut self, max: Option<u32>) -> Self {
193        self.max_subscriptions_per_connection = max;
194        self
195    }
196
197    /// Set remote subscription fields (federation passthrough).
198    ///
199    /// Maps subscription field names to the owning subgraph's `WebSocket` URL.
200    #[must_use]
201    pub fn with_remote_subscription_fields(mut self, fields: HashMap<String, String>) -> Self {
202        self.remote_subscription_fields = Arc::new(fields);
203        self
204    }
205}
206
207/// `WebSocket` upgrade handler for subscriptions.
208///
209/// Negotiates the `WebSocket` sub-protocol from the `Sec-WebSocket-Protocol`
210/// header. Supports `graphql-transport-ws` (modern) and `graphql-ws` (legacy).
211/// Defaults to `graphql-transport-ws` when no header is present.
212/// Returns `400 Bad Request` for unrecognised protocols.
213pub async fn subscription_handler(
214    headers: HeaderMap,
215    OptionalSecurityContext(security_context): OptionalSecurityContext,
216    ws: WebSocketUpgrade,
217    State(state): State<SubscriptionState>,
218) -> impl IntoResponse {
219    let protocol_header = headers.get("sec-websocket-protocol").and_then(|v| v.to_str().ok());
220
221    let protocol = match protocol_header {
222        None => WsProtocol::GraphqlTransportWs,
223        Some(header) => {
224            if let Some(p) = WsProtocol::from_header(Some(header)) {
225                p
226            } else {
227                warn!(header = %header, "Unknown WebSocket sub-protocol requested");
228                return axum::http::StatusCode::BAD_REQUEST.into_response();
229            }
230        },
231    };
232
233    // Resolve the tenant key exactly as the GraphQL handler does: the JWT
234    // `tenant_id` (trusted) takes precedence over the `X-Tenant-ID` header and
235    // the Host-domain lookup, and conflicting sources are rejected when strict
236    // validation is enabled. Previously this passed `None, None, false`,
237    // silently dropping JWT precedence, ignoring an installed domain registry,
238    // and disabling strict cross-source validation (#331).
239    let tenant_id = match resolve_subscription_tenant(security_context.as_ref(), &headers, &state) {
240        Ok(tenant_id) => tenant_id,
241        Err(e) => {
242            warn!(error = %e, "Subscription tenant resolution rejected the upgrade");
243            return axum::http::StatusCode::BAD_REQUEST.into_response();
244        },
245    };
246
247    ws.protocols([protocol.as_str()])
248        .on_upgrade(move |socket| {
249            handle_subscription_connection(socket, state, protocol, tenant_id, security_context)
250        })
251        .into_response()
252}
253
254/// Resolve the subscription's tenant key, mirroring the GraphQL handler's
255/// dispatch (`routes/graphql/handler.rs`): JWT `tenant_id` > `X-Tenant-ID`
256/// header > Host-domain registry, with cross-source conflict rejection when the
257/// schema has RLS configured (`strict_tenant_validation`).
258///
259/// # Errors
260///
261/// Returns `FraiseQLError::Validation` when the `X-Tenant-ID` header is invalid,
262/// or — under `strict_tenant_validation` — when the resolved sources disagree.
263fn resolve_subscription_tenant(
264    security_context: Option<&SecurityContext>,
265    headers: &HeaderMap,
266    state: &SubscriptionState,
267) -> fraiseql_error::Result<Option<String>> {
268    super::graphql::TenantKeyResolver::resolve(
269        security_context,
270        headers,
271        state.domain_registry.as_deref(),
272        state.strict_tenant_validation,
273    )
274}
275
276/// Handle a `WebSocket` subscription connection.
277#[allow(clippy::cognitive_complexity)] // Reason: WebSocket protocol state machine with message routing and lifecycle management
278async fn handle_subscription_connection(
279    socket: WebSocket,
280    state: SubscriptionState,
281    protocol: WsProtocol,
282    tenant_id: Option<String>,
283    principal: Option<SecurityContext>,
284) {
285    let connection_id = uuid::Uuid::new_v4().to_string();
286    let codec = ProtocolCodec::new(protocol);
287    info!(
288        connection_id = %connection_id,
289        protocol = %protocol.as_str(),
290        "WebSocket connection established"
291    );
292
293    let (mut sender, mut receiver) = socket.split();
294
295    // Wait for connection_init with timeout
296    let init_result = tokio::time::timeout(CONNECTION_INIT_TIMEOUT, async {
297        while let Some(msg) = receiver.next().await {
298            match msg {
299                Ok(Message::Text(text)) => {
300                    if let Ok(client_msg) = codec.decode(&text) {
301                        if client_msg.parsed_type() == Some(ClientMessageType::ConnectionInit) {
302                            return Some(client_msg);
303                        }
304                    }
305                },
306                Ok(Message::Close(_)) => return None,
307                Err(e) => {
308                    error!(error = %e, "WebSocket error during init");
309                    return None;
310                },
311                _ => {},
312            }
313        }
314        None
315    })
316    .await;
317
318    // Handle init timeout or failure
319    let _init_payload = match init_result {
320        Ok(Some(msg)) => {
321            // Call lifecycle on_connect hook
322            let params = msg.payload.clone().unwrap_or(serde_json::json!({}));
323            if let Err(reason) = state.lifecycle.on_connect(&params, &connection_id).await {
324                warn!(
325                    connection_id = %connection_id,
326                    reason = %reason,
327                    "Lifecycle on_connect rejected connection"
328                );
329                WS_CONNECTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
330                // Best-effort: connection is already being terminated.
331                let _ = sender
332                    .send(Message::Close(Some(axum::extract::ws::CloseFrame {
333                        code:   4400,
334                        reason: reason.into(),
335                    })))
336                    .await;
337                return;
338            }
339
340            // Send connection_ack
341            let ack = ServerMessage::connection_ack(None);
342            if let Err(send_err) = send_server_message(&codec, &mut sender, ack).await {
343                error!(connection_id = %connection_id, error = %send_err, "Failed to send connection_ack");
344                return;
345            }
346            WS_CONNECTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
347            info!(connection_id = %connection_id, "Connection initialized");
348            msg.payload
349        },
350        Ok(None) => {
351            warn!(connection_id = %connection_id, "Connection closed during init");
352            return;
353        },
354        Err(_) => {
355            warn!(connection_id = %connection_id, "Connection init timeout");
356            // Best-effort: connection is already being terminated.
357            let _ = sender
358                .send(Message::Close(Some(axum::extract::ws::CloseFrame {
359                    code:   CloseCode::ConnectionInitTimeout.code(),
360                    reason: CloseCode::ConnectionInitTimeout.reason().into(),
361                })))
362                .await;
363            return;
364        },
365    };
366
367    // Track active operations (operation_id -> subscription_id)
368    let mut active_operations: HashMap<String, SubscriptionId> = HashMap::new();
369
370    // Remote subscription message output channel.
371    //
372    // Forwarder tasks (federation feature) send pre-encoded ServerMessage values here.
373    // The channel is always present so the select! loop has a uniform branch regardless
374    // of whether any remote subscriptions are active.
375    let (remote_msg_tx, mut remote_msg_rx) = tokio::sync::mpsc::channel::<ServerMessage>(64);
376
377    // Subscribe to event broadcast
378    let mut event_receiver = state.manager.receiver();
379
380    // Ping/keepalive timer
381    let mut ping_interval = tokio::time::interval(PING_INTERVAL);
382    ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
383
384    // A44 — Token expiry re-check on long-lived subscriptions.
385    //
386    // JWTs validated at ConnectionInit may expire while the WebSocket is open.
387    // The check below should be added when the auth layer surfaces expiry data:
388    //
389    //   1. At ConnectionInit, extract the `exp` claim from the JWT and store it: `let
390    //      token_expires_at: Option<std::time::Instant> = extract_exp(&init_payload);`
391    //
392    //   2. In the select! loop (before processing each client message or broadcast event), check
393    //      expiry: ```rust,ignore if token_expires_at.is_some_and(|exp| std::time::Instant::now()
394    //      >= exp) { warn!(connection_id = %connection_id, "Token expired; closing WebSocket"); let
395    //      _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame { code:
396    //      CloseCode::Unauthorized.code(), reason: "Token expired".into(), }))).await; break; } ```
397    //
398    // This requires the lifecycle `on_connect` hook or the JWT middleware to return
399    // the expiry time, which is not yet threaded through `SubscriptionState`.
400    // Tracked as A44 in the remediation plan.
401
402    // Main message loop
403    loop {
404        tokio::select! {
405            msg = receiver.next() => {
406                match msg {
407                    Some(Ok(Message::Text(text))) => {
408                        if let Err(close_code) = handle_client_message(
409                            &text,
410                            &connection_id,
411                            &state,
412                            &codec,
413                            &mut active_operations,
414                            remote_msg_tx.clone(),
415                            &mut sender,
416                            tenant_id.as_deref(),
417                            principal.as_ref(),
418                        ).await {
419                            // Best-effort: connection is already being closed.
420                            let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
421                                code: close_code.code(),
422                                reason: close_code.reason().into(),
423                            }))).await;
424                            break;
425                        }
426                    }
427                    Some(Ok(Message::Ping(data))) => {
428                        // Best-effort: if the connection is already dead the pong will fail.
429                        let _ = sender.send(Message::Pong(data)).await;
430                    }
431                    Some(Ok(Message::Close(_))) => {
432                        info!(connection_id = %connection_id, "Client closed connection");
433                        break;
434                    }
435                    Some(Err(e)) => {
436                        error!(connection_id = %connection_id, error = %e, "WebSocket error");
437                        break;
438                    }
439                    None => {
440                        info!(connection_id = %connection_id, "WebSocket stream ended");
441                        break;
442                    }
443                    _ => {}
444                }
445            }
446
447            event = event_receiver.recv() => {
448                match event {
449                    Ok(payload) => {
450                        // Defense-in-depth tenant guard: when both the connection and the
451                        // event carry an explicit tenant_id they must agree. Primary
452                        // isolation is already guaranteed by subscription_id UUIDs, but
453                        // this check catches any future path that introduces deterministic
454                        // subscription IDs (which could collide across tenants).
455                        let tenant_matches = match (
456                            tenant_id.as_deref(),
457                            payload.event.tenant_id.as_deref(),
458                        ) {
459                            (Some(conn_tid), Some(evt_tid)) => conn_tid == evt_tid,
460                            _ => true, // either side absent → no conflict
461                        };
462                        if tenant_matches {
463                            if let Some((op_id, _)) = active_operations
464                                .iter()
465                                .find(|(_, sub_id)| **sub_id == payload.subscription_id)
466                            {
467                                let msg = create_next_message(op_id, &payload);
468                                if send_server_message(&codec, &mut sender, msg).await.is_err() {
469                                    warn!(connection_id = %connection_id, "Failed to send event");
470                                    break;
471                                }
472                            }
473                        }
474                    }
475                    Err(broadcast::error::RecvError::Lagged(n)) => {
476                        warn!(connection_id = %connection_id, lagged = n, "Event receiver lagged");
477                    }
478                    Err(broadcast::error::RecvError::Closed) => {
479                        error!(connection_id = %connection_id, "Event channel closed");
480                        break;
481                    }
482                }
483            }
484
485            remote_msg = remote_msg_rx.recv() => {
486                // Remote subscription event forwarded by a federation forwarder task.
487                // The channel is always present; it only carries messages when the
488                // federation feature is enabled and remote subscriptions are active.
489                if let Some(msg) = remote_msg {
490                    if send_server_message(&codec, &mut sender, msg).await.is_err() {
491                        warn!(connection_id = %connection_id, "Failed to send remote subscription message");
492                        break;
493                    }
494                }
495            }
496
497            _ = ping_interval.tick() => {
498                let msg = ServerMessage::ping(None);
499                if send_server_message(&codec, &mut sender, msg).await.is_err() {
500                    warn!(connection_id = %connection_id, "Failed to send ping/keepalive");
501                    break;
502                }
503            }
504        }
505    }
506
507    // Cleanup
508    state.manager.unsubscribe_connection(&connection_id);
509    state.lifecycle.on_disconnect(&connection_id).await;
510    info!(connection_id = %connection_id, "WebSocket connection closed");
511}
512
513/// Handle a client message.
514///
515/// Returns `Ok(())` on success, or `Err(CloseCode)` if the connection should be closed.
516///
517/// `remote_msg_tx` is used by federation forwarder tasks to send pre-encoded
518/// `ServerMessage` values back to the client connection loop.
519#[allow(clippy::cognitive_complexity)] // Reason: WebSocket message dispatch with subscribe/unsubscribe/query protocol handling
520#[allow(clippy::too_many_arguments)] // Reason: WebSocket handler needs connection state, protocol codec, and tenant context
521async fn handle_client_message(
522    text: &str,
523    connection_id: &str,
524    state: &SubscriptionState,
525    codec: &ProtocolCodec,
526    active_operations: &mut HashMap<String, SubscriptionId>,
527    remote_msg_tx: tokio::sync::mpsc::Sender<ServerMessage>,
528    sender: &mut futures::stream::SplitSink<WebSocket, Message>,
529    tenant_id: Option<&str>,
530    principal: Option<&SecurityContext>,
531) -> Result<(), CloseCode> {
532    // remote_msg_tx is only consumed inside the #[cfg(feature = "federation")] block below.
533    // When the feature is disabled the parameter goes unused; suppress the warning.
534    #[cfg(not(feature = "federation"))]
535    let _ = &remote_msg_tx;
536
537    let client_msg: ClientMessage = codec.decode(text).map_err(|e| {
538        warn!(error = %e, "Failed to parse client message");
539        CloseCode::ProtocolError
540    })?;
541
542    match client_msg.parsed_type() {
543        Some(ClientMessageType::Ping) => {
544            let pong = ServerMessage::pong(client_msg.payload);
545            // Best-effort: if the connection is already dead the pong will fail.
546            let _ = send_server_message(codec, sender, pong).await;
547        },
548
549        Some(ClientMessageType::Pong) => {
550            debug!(connection_id = %connection_id, "Received pong");
551        },
552
553        Some(ClientMessageType::Subscribe) => {
554            let payload: SubscribePayload = client_msg.subscription_payload().ok_or_else(|| {
555                warn!("Invalid subscribe payload");
556                CloseCode::ProtocolError
557            })?;
558
559            let op_id = client_msg.id.ok_or_else(|| {
560                warn!("Subscribe message missing operation ID");
561                CloseCode::ProtocolError
562            })?;
563
564            // Check for duplicate operation ID
565            if active_operations.contains_key(&op_id) {
566                warn!(operation_id = %op_id, "Duplicate operation ID");
567                return Err(CloseCode::SubscriberAlreadyExists);
568            }
569
570            // Enforce per-connection subscription limit
571            if let Some(max) = state.max_subscriptions_per_connection {
572                if active_operations.len() >= max as usize {
573                    warn!(
574                        connection_id = %connection_id,
575                        active = active_operations.len(),
576                        max = max,
577                        "Subscription limit reached"
578                    );
579                    WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
580                    let error = ServerMessage::error(
581                        &op_id,
582                        vec![GraphQLError::with_code(
583                            format!("Maximum subscriptions per connection ({max}) reached"),
584                            "SUBSCRIPTION_LIMIT_REACHED",
585                        )],
586                    );
587                    if let Err(e) = send_server_message(codec, sender, error).await {
588                        debug!(connection_id = %connection_id, error = %e, "Could not send subscription limit error to client");
589                    }
590                    return Ok(());
591                }
592            }
593
594            // Extract subscription name from query
595            let Some(subscription_name) = extract_subscription_name(&payload.query) else {
596                let error = ServerMessage::error(
597                    &op_id,
598                    vec![GraphQLError::with_code(
599                        "Could not parse subscription query",
600                        "PARSE_ERROR",
601                    )],
602                );
603                if let Err(e) = send_server_message(codec, sender, error).await {
604                    debug!(connection_id = %connection_id, error = %e, "Could not send parse error to client");
605                }
606                return Ok(());
607            };
608
609            // Call lifecycle on_subscribe hook
610            // HashMap<String, Value> serialization is infallible; the error path cannot occur.
611            let variables_value = serde_json::to_value(&payload.variables)
612                .expect("HashMap<String, serde_json::Value> serialization is infallible");
613
614            // #422: operation-level authorization at subscription establishment.
615            // The per-event delivery does not route through the executor, so the
616            // subscription is authorized once, here, with the connection's principal
617            // (or `None` when anonymous). Fail-closed: a `Deny` or any policy error
618            // rejects the subscription with a `FORBIDDEN` GraphQL-WS error.
619            if let Some(authorizer) = state.authorizer.as_ref() {
620                let ops = [(OperationKind::Subscription, subscription_name.clone())];
621                if let Err(err) =
622                    enforce_authz(authorizer.as_ref(), principal, &ops, Some(&variables_value))
623                {
624                    warn!(
625                        connection_id = %connection_id,
626                        subscription = %subscription_name,
627                        "Operation authorizer denied the subscription"
628                    );
629                    WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
630                    let error = ServerMessage::error(
631                        &op_id,
632                        vec![GraphQLError::with_code(err.to_string(), "FORBIDDEN")],
633                    );
634                    if let Err(e) = send_server_message(codec, sender, error).await {
635                        debug!(connection_id = %connection_id, error = %e, "Could not send authorization denial to client");
636                    }
637                    return Ok(());
638                }
639            }
640
641            if let Err(reason) = state
642                .lifecycle
643                .on_subscribe(&subscription_name, &variables_value, connection_id)
644                .await
645            {
646                warn!(
647                    connection_id = %connection_id,
648                    subscription = %subscription_name,
649                    reason = %reason,
650                    "Lifecycle on_subscribe rejected subscription"
651                );
652                WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
653                let error = ServerMessage::error(
654                    &op_id,
655                    vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REJECTED")],
656                );
657                if let Err(e) = send_server_message(codec, sender, error).await {
658                    debug!(connection_id = %connection_id, error = %e, "Could not send subscription rejection to client");
659                }
660                return Ok(());
661            }
662
663            // Forward to remote subgraph when the subscription field is owned remotely.
664            #[cfg(feature = "federation")]
665            if let Some(subgraph_url) = state.remote_subscription_fields.get(&subscription_name) {
666                use fraiseql_federation::subscription_forwarder::{
667                    ForwardedEvent, SubscriptionForwarder,
668                };
669
670                match SubscriptionForwarder::new(subgraph_url) {
671                    Ok(forwarder) => {
672                        // Create a channel so the forwarder task can send us raw events.
673                        let (event_tx, mut event_rx) =
674                            tokio::sync::mpsc::channel::<ForwardedEvent>(32);
675
676                        // Task 1: run the WebSocket forwarder (sends ForwardedEvent to event_tx).
677                        let fwd_op = op_id.clone();
678                        let fwd_query = payload.query.clone();
679                        let fwd_vars = variables_value.clone();
680                        tokio::spawn(async move {
681                            if let Err(e) =
682                                forwarder.forward(&fwd_op, &fwd_query, fwd_vars, event_tx).await
683                            {
684                                warn!(error = %e, "Remote subscription forwarder failed");
685                            }
686                        });
687
688                        // Task 2: relay ForwardedEvent → ServerMessage → client.
689                        let relay_op = op_id.clone();
690                        let relay_tx = remote_msg_tx.clone();
691                        tokio::spawn(async move {
692                            while let Some(event) = event_rx.recv().await {
693                                let server_msg = match event {
694                                    ForwardedEvent::Next(data) => {
695                                        ServerMessage::next(&relay_op, data)
696                                    },
697                                    ForwardedEvent::Error(errors) => {
698                                        let errors_vec = errors.as_array().map_or_else(
699                                            || {
700                                                vec![GraphQLError::with_code(
701                                                    errors.to_string(),
702                                                    "REMOTE_ERROR",
703                                                )]
704                                            },
705                                            |arr| {
706                                                arr.iter()
707                                                    .map(|e| {
708                                                        GraphQLError::with_code(
709                                                            e.get("message")
710                                                                .and_then(|v| v.as_str())
711                                                                .unwrap_or("Remote subgraph error"),
712                                                            "REMOTE_ERROR",
713                                                        )
714                                                    })
715                                                    .collect()
716                                            },
717                                        );
718                                        ServerMessage::error(&relay_op, errors_vec)
719                                    },
720                                    ForwardedEvent::Complete => ServerMessage::complete(&relay_op),
721                                };
722                                if relay_tx.send(server_msg).await.is_err() {
723                                    break; // Client disconnected
724                                }
725                            }
726                        });
727
728                        WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
729                        info!(
730                            connection_id = %connection_id,
731                            operation_id = %op_id,
732                            subscription = %subscription_name,
733                            "Subscription forwarded to remote subgraph"
734                        );
735                        return Ok(());
736                    },
737                    Err(e) => {
738                        let error = ServerMessage::error(
739                            &op_id,
740                            vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
741                        );
742                        if let Err(send_err) = send_server_message(codec, sender, error).await {
743                            debug!(connection_id = %connection_id, error = %send_err, "Could not send forwarding error to client");
744                        }
745                        return Ok(());
746                    },
747                }
748            }
749
750            // Validate client-provided tenant variable against server-resolved
751            if let Some(server_tid) = tenant_id {
752                if let Some(client_tid) = variables_value.get("tenant_id").and_then(|v| v.as_str())
753                {
754                    if client_tid != server_tid {
755                        let error = ServerMessage::error(
756                            &op_id,
757                            vec![GraphQLError::with_code(
758                                format!(
759                                    "Tenant mismatch: client provided '{client_tid}', server resolved '{server_tid}'"
760                                ),
761                                "TENANT_MISMATCH",
762                            )],
763                        );
764                        if let Err(send_err) = send_server_message(codec, sender, error).await {
765                            debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant mismatch error to client");
766                        }
767                        return Ok(());
768                    }
769                }
770            }
771
772            // Build context with server-resolved tenant_id
773            let mut context = serde_json::json!({});
774            if let Some(tid) = tenant_id {
775                context["tenant_id"] = serde_json::Value::String(tid.to_string());
776            }
777
778            // Subscribe locally (field is owned by this subgraph)
779            match state.manager.subscribe(
780                &subscription_name,
781                context,
782                variables_value,
783                connection_id,
784            ) {
785                Ok(sub_id) => {
786                    active_operations.insert(op_id.clone(), sub_id);
787                    WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
788                    info!(
789                        connection_id = %connection_id,
790                        operation_id = %op_id,
791                        subscription = %subscription_name,
792                        "Subscription started"
793                    );
794                },
795                Err(e) => {
796                    let error = ServerMessage::error(
797                        &op_id,
798                        vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
799                    );
800                    if let Err(send_err) = send_server_message(codec, sender, error).await {
801                        debug!(connection_id = %connection_id, error = %send_err, "Could not send subscription error to client");
802                    }
803                },
804            }
805        },
806
807        Some(ClientMessageType::Complete) => {
808            let op_id = client_msg.id.ok_or_else(|| {
809                warn!("Complete message missing operation ID");
810                CloseCode::ProtocolError
811            })?;
812
813            if let Some(sub_id) = active_operations.remove(&op_id) {
814                if let Err(e) = state.manager.unsubscribe(sub_id) {
815                    warn!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe; subscription may be leaked");
816                }
817                state.lifecycle.on_unsubscribe(&op_id, connection_id).await;
818                info!(
819                    connection_id = %connection_id,
820                    operation_id = %op_id,
821                    "Subscription completed"
822                );
823            }
824        },
825
826        Some(ClientMessageType::ConnectionInit) => {
827            warn!(connection_id = %connection_id, "Duplicate connection_init");
828            return Err(CloseCode::TooManyInitRequests);
829        },
830
831        None => {
832            warn!(message_type = %client_msg.message_type, "Unknown message type");
833        },
834        // Reason: non_exhaustive requires catch-all for cross-crate matches
835        _ => {
836            warn!(message_type = %client_msg.message_type, "Unrecognized message type");
837        },
838    }
839
840    Ok(())
841}
842
843/// Send a server message through the codec, handling protocol translation.
844async fn send_server_message(
845    codec: &ProtocolCodec,
846    sender: &mut futures::stream::SplitSink<WebSocket, Message>,
847    msg: ServerMessage,
848) -> Result<(), String> {
849    match codec.encode(&msg) {
850        Ok(Some(json)) => sender.send(Message::Text(json.into())).await.map_err(|e| e.to_string()),
851        Ok(None) => Ok(()), // Message suppressed by codec (e.g. pong in legacy mode)
852        Err(e) => Err(e.to_string()),
853    }
854}
855
856/// Create a "next" message for a subscription event.
857fn create_next_message(operation_id: &str, payload: &SubscriptionPayload) -> ServerMessage {
858    let data = serde_json::json!({
859        payload.subscription_name.clone(): payload.data
860    });
861    ServerMessage::next(operation_id, data)
862}
863
864/// Extract subscription name from a GraphQL subscription query.
865pub(crate) fn extract_subscription_name(query: &str) -> Option<String> {
866    let query = query.trim();
867
868    let sub_idx = query.find("subscription")?;
869    let after_sub = &query[sub_idx + "subscription".len()..];
870
871    let brace_idx = after_sub.find('{')?;
872    let after_brace = after_sub[brace_idx + 1..].trim_start();
873
874    let name_end = after_brace
875        .find(|c: char| !c.is_alphanumeric() && c != '_')
876        .unwrap_or(after_brace.len());
877
878    if name_end == 0 {
879        return None;
880    }
881
882    Some(after_brace[..name_end].to_string())
883}
884
885#[cfg(test)]
886mod tests;