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    schema::{CompiledSchema, OwnerCondition, SubscriptionPolicy},
52    security::{Authorizer, OperationKind, SecurityContext, authorizer::enforce_authz},
53};
54use futures::{SinkExt, StreamExt};
55use tokio::sync::broadcast;
56use tracing::{debug, error, info, warn};
57
58use crate::{
59    extractors::OptionalSecurityContext,
60    routes::graphql::{DomainRegistry, TenantStatusSource},
61    subscriptions::{
62        lifecycle::SubscriptionLifecycle,
63        protocol::{ProtocolCodec, WsProtocol},
64    },
65};
66
67// ── Subscription metrics (module-level atomics) ──────────────────────
68
69static WS_CONNECTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
70static WS_CONNECTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
71static WS_SUBSCRIPTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
72static WS_SUBSCRIPTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
73
74/// Subscription metrics for Prometheus export.
75#[must_use]
76pub fn subscription_metrics() -> SubscriptionMetrics {
77    SubscriptionMetrics {
78        connections_accepted:   WS_CONNECTIONS_ACCEPTED.load(Ordering::Relaxed),
79        connections_rejected:   WS_CONNECTIONS_REJECTED.load(Ordering::Relaxed),
80        subscriptions_accepted: WS_SUBSCRIPTIONS_ACCEPTED.load(Ordering::Relaxed),
81        subscriptions_rejected: WS_SUBSCRIPTIONS_REJECTED.load(Ordering::Relaxed),
82    }
83}
84
85/// Reset all subscription counters to zero.
86///
87/// Call this at the start of each test that checks counter values to avoid
88/// cross-test interference from the module-level statics.
89#[cfg(test)]
90pub fn reset_metrics_for_test() {
91    WS_CONNECTIONS_ACCEPTED.store(0, Ordering::SeqCst);
92    WS_CONNECTIONS_REJECTED.store(0, Ordering::SeqCst);
93    WS_SUBSCRIPTIONS_ACCEPTED.store(0, Ordering::SeqCst);
94    WS_SUBSCRIPTIONS_REJECTED.store(0, Ordering::SeqCst);
95}
96
97/// Snapshot of subscription counters.
98pub struct SubscriptionMetrics {
99    /// Total `WebSocket` connections accepted (after `on_connect`).
100    pub connections_accepted:   u64,
101    /// Total `WebSocket` connections rejected by lifecycle hook.
102    pub connections_rejected:   u64,
103    /// Total subscriptions accepted (after `on_subscribe`).
104    pub subscriptions_accepted: u64,
105    /// Total subscriptions rejected (by hook or limit).
106    pub subscriptions_rejected: u64,
107}
108
109/// Connection initialization timeout (5 seconds per graphql-ws spec).
110const CONNECTION_INIT_TIMEOUT: Duration = Duration::from_secs(5);
111
112/// Ping/keepalive interval.
113const PING_INTERVAL: Duration = Duration::from_secs(30);
114
115/// State for subscription `WebSocket` handler.
116#[derive(Clone)]
117pub struct SubscriptionState {
118    /// Subscription manager.
119    pub manager: Arc<SubscriptionManager>,
120    /// Lifecycle hooks.
121    pub lifecycle: Arc<dyn SubscriptionLifecycle>,
122    /// Maximum subscriptions per connection (`None` = unlimited).
123    pub max_subscriptions_per_connection: Option<u32>,
124    /// Subscription fields owned by remote subgraphs.
125    ///
126    /// Maps root subscription field name to the subgraph `WebSocket` URL.
127    /// Empty when federation is disabled or no remote subscription fields are declared.
128    pub remote_subscription_fields: Arc<HashMap<String, String>>,
129    /// Host-header → tenant-key domain registry. `None` until a host binary
130    /// installs one (mirrors `AppState::domain_registry`).
131    pub domain_registry: Option<Arc<DomainRegistry>>,
132    /// Reject conflicting tenant sources (JWT vs `X-Tenant-ID` vs Host) on the
133    /// upgrade. Driven by `schema.has_rls_configured()`, mirroring the GraphQL
134    /// handler's strict tenant validation.
135    pub strict_tenant_validation: bool,
136    /// Optional operation-level authorizer (#422). When set, each subscription is
137    /// authorized at establishment with [`OperationKind::Subscription`], the
138    /// subscription field name, and the connection's principal. `None` until a host
139    /// binary installs one (from `Executor::config().authorizer`).
140    pub authorizer: Option<Arc<dyn Authorizer>>,
141    /// Optional tenant-status source (M-tenant-ws-suspended). When set, a new
142    /// subscription whose resolved tenant is suspended is rejected, and event
143    /// delivery to a connection whose tenant is suspended is paused. `None` until
144    /// a host binary installs a multi-tenant registry.
145    pub tenant_status_source: Option<Arc<dyn TenantStatusSource>>,
146    /// Per-subscription-field row-visibility policies (#596), keyed by subscription
147    /// field name, resolved at mount time from the target entity's compiled
148    /// `subscription_policy`. A subscription named here derives a **server-owned** owner
149    /// condition from the connection's enriched identity at subscribe time — fail-closed
150    /// when the identity is unresolvable. A subscription with no policy keeps today's
151    /// behavior (no back-compat break).
152    pub subscription_policies: Arc<HashMap<String, SubscriptionPolicy>>,
153    /// Enriched-identity resolver (#539). When set, the connection's `SecurityContext`
154    /// is enriched at subscribe time (only for policy-declaring subscriptions) so the
155    /// `fraiseql.enriched.*` owner field is server-resolved, never client-asserted.
156    /// `None` disables enrichment — a policy-declaring subscription then fails closed.
157    #[cfg(feature = "auth")]
158    pub identity_resolver: Option<Arc<crate::identity::IdentityResolver>>,
159    /// Service-account authenticator (ADR-0018). Lets a daemon authenticate the `/ws`
160    /// upgrade with its secret on the api-key header — the same seam the GraphQL path
161    /// uses — so a service principal can hold a policy-scoped subscription.
162    pub service_account_authenticator:
163        Option<Arc<crate::service_account::ServiceAccountAuthenticator>>,
164}
165
166impl SubscriptionState {
167    /// Create new subscription state.
168    pub fn new(manager: Arc<SubscriptionManager>) -> Self {
169        Self {
170            manager,
171            lifecycle: Arc::new(crate::subscriptions::lifecycle::NoopLifecycle),
172            max_subscriptions_per_connection: None,
173            remote_subscription_fields: Arc::new(HashMap::new()),
174            domain_registry: None,
175            strict_tenant_validation: false,
176            authorizer: None,
177            tenant_status_source: None,
178            subscription_policies: Arc::new(HashMap::new()),
179            #[cfg(feature = "auth")]
180            identity_resolver: None,
181            service_account_authenticator: None,
182        }
183    }
184
185    /// Install the per-subscription row-visibility policies (#596). Typically built by
186    /// [`build_subscription_policies`] from the compiled schema at mount time.
187    #[must_use]
188    pub fn with_subscription_policies(
189        mut self,
190        policies: Arc<HashMap<String, SubscriptionPolicy>>,
191    ) -> Self {
192        self.subscription_policies = policies;
193        self
194    }
195
196    /// Install the service-account authenticator (ADR-0018) so a daemon can authenticate
197    /// the `/ws` upgrade with its secret on the api-key header.
198    #[must_use]
199    pub fn with_service_account_authenticator(
200        mut self,
201        authenticator: Option<Arc<crate::service_account::ServiceAccountAuthenticator>>,
202    ) -> Self {
203        self.service_account_authenticator = authenticator;
204        self
205    }
206
207    /// Install the enriched-identity resolver (#539) used to derive row-visibility owner
208    /// boundaries at subscribe time. `None` leaves policy-declaring subscriptions
209    /// fail-closed (refused).
210    #[cfg(feature = "auth")]
211    #[must_use]
212    pub fn with_identity_resolver(
213        mut self,
214        resolver: Option<Arc<crate::identity::IdentityResolver>>,
215    ) -> Self {
216        self.identity_resolver = resolver;
217        self
218    }
219
220    /// Install the tenant-status source (M-tenant-ws-suspended). When set, new
221    /// subscriptions for a suspended tenant are rejected and event delivery to a
222    /// suspended tenant is paused. Typically the `TenantExecutorRegistry`.
223    #[must_use]
224    pub fn with_tenant_status_source(
225        mut self,
226        source: Option<Arc<dyn TenantStatusSource>>,
227    ) -> Self {
228        self.tenant_status_source = source;
229        self
230    }
231
232    /// Install the operation-level authorizer (#422). When set, every subscription is
233    /// authorized at establishment; a `Deny` (or any policy error) rejects the
234    /// subscription with a `FORBIDDEN` GraphQL-WS error. Typically populated from
235    /// `Executor::config().authorizer`.
236    #[must_use]
237    pub fn with_authorizer(mut self, authorizer: Option<Arc<dyn Authorizer>>) -> Self {
238        self.authorizer = authorizer;
239        self
240    }
241
242    /// Install the tenant-resolution context — the Host-header domain registry
243    /// and strict-validation flag — so the subscription upgrade dispatches the
244    /// tenant key the same way the GraphQL handler does (JWT `tenant_id` >
245    /// `X-Tenant-ID` header > Host domain, with cross-source conflict rejection
246    /// when `strict_tenant_validation` is set). See
247    /// [`crate::routes::graphql::TenantKeyResolver`].
248    #[must_use]
249    pub fn with_tenant_context(
250        mut self,
251        domain_registry: Arc<DomainRegistry>,
252        strict_tenant_validation: bool,
253    ) -> Self {
254        self.domain_registry = Some(domain_registry);
255        self.strict_tenant_validation = strict_tenant_validation;
256        self
257    }
258
259    /// Set lifecycle hooks.
260    #[must_use]
261    pub fn with_lifecycle(mut self, lifecycle: Arc<dyn SubscriptionLifecycle>) -> Self {
262        self.lifecycle = lifecycle;
263        self
264    }
265
266    /// Set maximum subscriptions per connection.
267    #[must_use]
268    pub const fn with_max_subscriptions(mut self, max: Option<u32>) -> Self {
269        self.max_subscriptions_per_connection = max;
270        self
271    }
272
273    /// Set remote subscription fields (federation passthrough).
274    ///
275    /// Maps subscription field names to the owning subgraph's `WebSocket` URL.
276    #[must_use]
277    pub fn with_remote_subscription_fields(mut self, fields: HashMap<String, String>) -> Self {
278        self.remote_subscription_fields = Arc::new(fields);
279        self
280    }
281}
282
283/// Build the per-subscription row-visibility policy map (#596) from the compiled schema.
284///
285/// For each subscription field, resolve its target entity's `subscription_policy` (if
286/// any) and key it by the subscription field name. The `/ws` handler consults this map
287/// at subscribe time; an entry means "derive a server-owned owner condition and refuse
288/// if the identity is unresolvable", absence means "unchanged behavior".
289#[must_use]
290pub fn build_subscription_policies(schema: &CompiledSchema) -> HashMap<String, SubscriptionPolicy> {
291    let mut policies = HashMap::new();
292    for sub in &schema.subscriptions {
293        if let Some(type_def) = schema.types.iter().find(|t| t.name.as_str() == sub.return_type) {
294            if let Some(policy) = &type_def.subscription_policy {
295                policies.insert(sub.name.clone(), policy.clone());
296            }
297        }
298    }
299    policies
300}
301
302/// Resolve the server-owned RLS conditions to enforce for a subscribe request (#596).
303///
304/// - `Ok(vec![])` — the subscription declares no policy (back-compat) **or** the principal holds a
305///   bypass role (full visibility);
306/// - `Ok(vec![(owner_field, identity_value)])` — a scoped subscriber sees only its rows;
307/// - `Err(reason)` — the policy applies but the identity is unresolvable; the caller **refuses**
308///   the subscription (fail-closed, never deliver-all).
309async fn resolve_subscription_rls(
310    state: &SubscriptionState,
311    subscription_name: &str,
312    principal: Option<&SecurityContext>,
313) -> Result<Vec<(String, serde_json::Value)>, String> {
314    let Some(policy) = state.subscription_policies.get(subscription_name) else {
315        return Ok(Vec::new());
316    };
317
318    // Enrich a copy of the principal so the owner boundary derives from the
319    // server-resolved `fraiseql.enriched.*` namespace, never a client claim. A failed
320    // or absent enrichment leaves the enriched field unresolvable → the derivation
321    // refuses below (fail-closed). Enrichment runs only for policy-declaring
322    // subscriptions, so unscoped streams pay nothing.
323    #[cfg(feature = "auth")]
324    let enriched = enrich_principal(state, principal).await;
325    #[cfg(feature = "auth")]
326    let effective = enriched.as_ref().or(principal);
327    #[cfg(not(feature = "auth"))]
328    let effective = principal;
329
330    derive_policy_conditions(policy, effective)
331}
332
333/// Enrich a clone of the principal's `SecurityContext` (#539). Best-effort: on a denial
334/// or resolver outage the enriched fields are simply absent, so a policy-declaring
335/// subscription fails closed at derivation. Returns `None` when there is no principal
336/// or no resolver configured.
337#[cfg(feature = "auth")]
338async fn enrich_principal(
339    state: &SubscriptionState,
340    principal: Option<&SecurityContext>,
341) -> Option<SecurityContext> {
342    let resolver = state.identity_resolver.as_ref()?;
343    let mut ctx = principal?.clone();
344    let _ = crate::identity::enrich_security_context(resolver, &mut ctx).await;
345    Some(ctx)
346}
347
348/// Adapt a policy's seam-neutral [`OwnerCondition`] to the `/ws` seam's `(field, value)`
349/// condition channel (#596). Fail-closed: `Refuse` → `Err`. A `None` principal (or one
350/// lacking the enriched field) derives to `Refuse`, so anonymous subscribers cannot see
351/// a policy-scoped entity's rows.
352fn derive_policy_conditions(
353    policy: &SubscriptionPolicy,
354    principal: Option<&SecurityContext>,
355) -> Result<Vec<(String, serde_json::Value)>, String> {
356    let empty = HashMap::new();
357    let attributes = principal.map_or(&empty, |ctx| &ctx.attributes);
358    let roles: &[String] = principal.map_or(&[], |ctx| ctx.roles.as_slice());
359    match policy.derive(attributes, roles) {
360        OwnerCondition::Bypass => Ok(Vec::new()),
361        OwnerCondition::Eq { field, value } => Ok(vec![(field, value)]),
362        OwnerCondition::Refuse(reason) => Err(reason),
363    }
364}
365
366/// `WebSocket` upgrade handler for subscriptions.
367///
368/// Negotiates the `WebSocket` sub-protocol from the `Sec-WebSocket-Protocol`
369/// header. Supports `graphql-transport-ws` (modern) and `graphql-ws` (legacy).
370/// Defaults to `graphql-transport-ws` when no header is present.
371/// Returns `400 Bad Request` for unrecognised protocols.
372pub async fn subscription_handler(
373    headers: HeaderMap,
374    OptionalSecurityContext(security_context): OptionalSecurityContext,
375    ws: WebSocketUpgrade,
376    State(state): State<SubscriptionState>,
377) -> impl IntoResponse {
378    let protocol_header = headers.get("sec-websocket-protocol").and_then(|v| v.to_str().ok());
379
380    let protocol = match protocol_header {
381        None => WsProtocol::GraphqlTransportWs,
382        Some(header) => {
383            if let Some(p) = WsProtocol::from_header(Some(header)) {
384                p
385            } else {
386                warn!(header = %header, "Unknown WebSocket sub-protocol requested");
387                return axum::http::StatusCode::BAD_REQUEST.into_response();
388            }
389        },
390    };
391
392    // ADR-0018: authenticate a service account presenting its secret on the api-key
393    // header — the same seam the GraphQL path uses — so a service principal can hold a
394    // policy-scoped subscription. A JWT principal AND a secret on one upgrade is
395    // ambiguous (#602) → 401; a present-but-unmatched secret → 401 (indistinguishable
396    // from an unknown account, no oracle).
397    let mut security_context = security_context;
398    if let Some(sa_auth) = state.service_account_authenticator.as_ref() {
399        match sa_auth.resolve(&headers, security_context.is_some()) {
400            crate::service_account::SaAuth::NoSecret => {},
401            crate::service_account::SaAuth::Authenticated(ctx) => security_context = Some(*ctx),
402            crate::service_account::SaAuth::Ambiguous
403            | crate::service_account::SaAuth::Unmatched => {
404                warn!(
405                    "Subscription upgrade rejected: ambiguous or unmatched service-account secret"
406                );
407                return axum::http::StatusCode::UNAUTHORIZED.into_response();
408            },
409        }
410    }
411
412    // Resolve the tenant key exactly as the GraphQL handler does: the JWT
413    // `tenant_id` (trusted) takes precedence over the `X-Tenant-ID` header and
414    // the Host-domain lookup, and conflicting sources are rejected when strict
415    // validation is enabled. Previously this passed `None, None, false`,
416    // silently dropping JWT precedence, ignoring an installed domain registry,
417    // and disabling strict cross-source validation (#331).
418    let tenant_id = match resolve_subscription_tenant(security_context.as_ref(), &headers, &state) {
419        Ok(tenant_id) => tenant_id,
420        Err(e) => {
421            warn!(error = %e, "Subscription tenant resolution rejected the upgrade");
422            return axum::http::StatusCode::BAD_REQUEST.into_response();
423        },
424    };
425
426    ws.protocols([protocol.as_str()])
427        .on_upgrade(move |socket| {
428            handle_subscription_connection(socket, state, protocol, tenant_id, security_context)
429        })
430        .into_response()
431}
432
433/// Resolve the subscription's tenant key, mirroring the GraphQL handler's
434/// dispatch (`routes/graphql/handler.rs`): JWT `tenant_id` > `X-Tenant-ID`
435/// header > Host-domain registry, with cross-source conflict rejection when the
436/// schema has RLS configured (`strict_tenant_validation`).
437///
438/// # Errors
439///
440/// Returns `FraiseQLError::Validation` when the `X-Tenant-ID` header is invalid,
441/// or — under `strict_tenant_validation` — when the resolved sources disagree.
442fn resolve_subscription_tenant(
443    security_context: Option<&SecurityContext>,
444    headers: &HeaderMap,
445    state: &SubscriptionState,
446) -> fraiseql_error::Result<Option<String>> {
447    super::graphql::TenantKeyResolver::resolve(
448        security_context,
449        headers,
450        state.domain_registry.as_deref(),
451        state.strict_tenant_validation,
452    )
453}
454
455/// Handle a `WebSocket` subscription connection.
456#[allow(clippy::cognitive_complexity)] // Reason: WebSocket protocol state machine with message routing and lifecycle management
457async fn handle_subscription_connection(
458    socket: WebSocket,
459    state: SubscriptionState,
460    protocol: WsProtocol,
461    tenant_id: Option<String>,
462    principal: Option<SecurityContext>,
463) {
464    let connection_id = uuid::Uuid::new_v4().to_string();
465    let codec = ProtocolCodec::new(protocol);
466    info!(
467        connection_id = %connection_id,
468        protocol = %protocol.as_str(),
469        "WebSocket connection established"
470    );
471
472    let (mut sender, mut receiver) = socket.split();
473
474    // Wait for connection_init with timeout
475    let init_result = tokio::time::timeout(CONNECTION_INIT_TIMEOUT, async {
476        while let Some(msg) = receiver.next().await {
477            match msg {
478                Ok(Message::Text(text)) => {
479                    if let Ok(client_msg) = codec.decode(&text) {
480                        if client_msg.parsed_type() == Some(ClientMessageType::ConnectionInit) {
481                            return Some(client_msg);
482                        }
483                    }
484                },
485                Ok(Message::Close(_)) => return None,
486                Err(e) => {
487                    error!(error = %e, "WebSocket error during init");
488                    return None;
489                },
490                _ => {},
491            }
492        }
493        None
494    })
495    .await;
496
497    // Handle init timeout or failure
498    let _init_payload = match init_result {
499        Ok(Some(msg)) => {
500            // Call lifecycle on_connect hook
501            let params = msg.payload.clone().unwrap_or(serde_json::json!({}));
502            if let Err(reason) = state.lifecycle.on_connect(&params, &connection_id).await {
503                warn!(
504                    connection_id = %connection_id,
505                    reason = %reason,
506                    "Lifecycle on_connect rejected connection"
507                );
508                WS_CONNECTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
509                // Best-effort: connection is already being terminated.
510                let _ = sender
511                    .send(Message::Close(Some(axum::extract::ws::CloseFrame {
512                        code:   4400,
513                        reason: reason.into(),
514                    })))
515                    .await;
516                return;
517            }
518
519            // Send connection_ack
520            let ack = ServerMessage::connection_ack(None);
521            if let Err(send_err) = send_server_message(&codec, &mut sender, ack).await {
522                error!(connection_id = %connection_id, error = %send_err, "Failed to send connection_ack");
523                return;
524            }
525            WS_CONNECTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
526            info!(connection_id = %connection_id, "Connection initialized");
527            msg.payload
528        },
529        Ok(None) => {
530            warn!(connection_id = %connection_id, "Connection closed during init");
531            return;
532        },
533        Err(_) => {
534            warn!(connection_id = %connection_id, "Connection init timeout");
535            // Best-effort: connection is already being terminated.
536            let _ = sender
537                .send(Message::Close(Some(axum::extract::ws::CloseFrame {
538                    code:   CloseCode::ConnectionInitTimeout.code(),
539                    reason: CloseCode::ConnectionInitTimeout.reason().into(),
540                })))
541                .await;
542            return;
543        },
544    };
545
546    // Track active operations (operation_id -> subscription_id)
547    let mut active_operations: HashMap<String, SubscriptionId> = HashMap::new();
548
549    // Remote subscription message output channel.
550    //
551    // Forwarder tasks (federation feature) send pre-encoded ServerMessage values here.
552    // The channel is always present so the select! loop has a uniform branch regardless
553    // of whether any remote subscriptions are active.
554    let (remote_msg_tx, mut remote_msg_rx) = tokio::sync::mpsc::channel::<ServerMessage>(64);
555
556    // Subscribe to event broadcast
557    let mut event_receiver = state.manager.receiver();
558
559    // Ping/keepalive timer
560    let mut ping_interval = tokio::time::interval(PING_INTERVAL);
561    ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
562
563    // A44 — Token expiry re-check on long-lived subscriptions.
564    //
565    // JWTs validated at ConnectionInit may expire while the WebSocket is open.
566    // The check below should be added when the auth layer surfaces expiry data:
567    //
568    //   1. At ConnectionInit, extract the `exp` claim from the JWT and store it: `let
569    //      token_expires_at: Option<std::time::Instant> = extract_exp(&init_payload);`
570    //
571    //   2. In the select! loop (before processing each client message or broadcast event), check
572    //      expiry: ```rust,ignore if token_expires_at.is_some_and(|exp| std::time::Instant::now()
573    //      >= exp) { warn!(connection_id = %connection_id, "Token expired; closing WebSocket"); let
574    //      _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame { code:
575    //      CloseCode::Unauthorized.code(), reason: "Token expired".into(), }))).await; break; } ```
576    //
577    // This requires the lifecycle `on_connect` hook or the JWT middleware to return
578    // the expiry time, which is not yet threaded through `SubscriptionState`.
579    // Tracked as A44 in the remediation plan.
580
581    // Main message loop
582    loop {
583        tokio::select! {
584            msg = receiver.next() => {
585                match msg {
586                    Some(Ok(Message::Text(text))) => {
587                        if let Err(close_code) = handle_client_message(
588                            &text,
589                            &connection_id,
590                            &state,
591                            &codec,
592                            &mut active_operations,
593                            remote_msg_tx.clone(),
594                            &mut sender,
595                            tenant_id.as_deref(),
596                            principal.as_ref(),
597                        ).await {
598                            // Best-effort: connection is already being closed.
599                            let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
600                                code: close_code.code(),
601                                reason: close_code.reason().into(),
602                            }))).await;
603                            break;
604                        }
605                    }
606                    Some(Ok(Message::Ping(data))) => {
607                        // Best-effort: if the connection is already dead the pong will fail.
608                        let _ = sender.send(Message::Pong(data)).await;
609                    }
610                    Some(Ok(Message::Close(_))) => {
611                        info!(connection_id = %connection_id, "Client closed connection");
612                        break;
613                    }
614                    Some(Err(e)) => {
615                        error!(connection_id = %connection_id, error = %e, "WebSocket error");
616                        break;
617                    }
618                    None => {
619                        info!(connection_id = %connection_id, "WebSocket stream ended");
620                        break;
621                    }
622                    _ => {}
623                }
624            }
625
626            event = event_receiver.recv() => {
627                match event {
628                    Ok(payload) => {
629                        // Defense-in-depth tenant guard: when both the connection and the
630                        // event carry an explicit tenant_id they must agree. Primary
631                        // isolation is already guaranteed by subscription_id UUIDs, but
632                        // this check catches any future path that introduces deterministic
633                        // subscription IDs (which could collide across tenants).
634                        let tenant_matches = match (
635                            tenant_id.as_deref(),
636                            payload.event.tenant_id.as_deref(),
637                        ) {
638                            (Some(conn_tid), Some(evt_tid)) => conn_tid == evt_tid,
639                            _ => true, // either side absent → no conflict
640                        };
641                        // M-tenant-ws-suspended: pause delivery while this
642                        // connection's tenant is suspended (re-checked per event so a
643                        // suspension mid-stream stops further delivery).
644                        let tenant_active = match (
645                            tenant_id.as_deref(),
646                            state.tenant_status_source.as_ref(),
647                        ) {
648                            (Some(tid), Some(src)) => !src.is_suspended(tid),
649                            _ => true,
650                        };
651                        if tenant_matches && tenant_active {
652                            if let Some((op_id, _)) = active_operations
653                                .iter()
654                                .find(|(_, sub_id)| **sub_id == payload.subscription_id)
655                            {
656                                let msg = create_next_message(op_id, &payload);
657                                if send_server_message(&codec, &mut sender, msg).await.is_err() {
658                                    warn!(connection_id = %connection_id, "Failed to send event");
659                                    break;
660                                }
661                            }
662                        }
663                    }
664                    Err(broadcast::error::RecvError::Lagged(n)) => {
665                        warn!(connection_id = %connection_id, lagged = n, "Event receiver lagged");
666                    }
667                    Err(broadcast::error::RecvError::Closed) => {
668                        error!(connection_id = %connection_id, "Event channel closed");
669                        break;
670                    }
671                }
672            }
673
674            remote_msg = remote_msg_rx.recv() => {
675                // Remote subscription event forwarded by a federation forwarder task.
676                // The channel is always present; it only carries messages when the
677                // federation feature is enabled and remote subscriptions are active.
678                if let Some(msg) = remote_msg {
679                    if send_server_message(&codec, &mut sender, msg).await.is_err() {
680                        warn!(connection_id = %connection_id, "Failed to send remote subscription message");
681                        break;
682                    }
683                }
684            }
685
686            _ = ping_interval.tick() => {
687                let msg = ServerMessage::ping(None);
688                if send_server_message(&codec, &mut sender, msg).await.is_err() {
689                    warn!(connection_id = %connection_id, "Failed to send ping/keepalive");
690                    break;
691                }
692            }
693        }
694    }
695
696    // Cleanup
697    state.manager.unsubscribe_connection(&connection_id);
698    state.lifecycle.on_disconnect(&connection_id).await;
699    info!(connection_id = %connection_id, "WebSocket connection closed");
700}
701
702/// Handle a client message.
703///
704/// Returns `Ok(())` on success, or `Err(CloseCode)` if the connection should be closed.
705///
706/// `remote_msg_tx` is used by federation forwarder tasks to send pre-encoded
707/// `ServerMessage` values back to the client connection loop.
708#[allow(clippy::cognitive_complexity)] // Reason: WebSocket message dispatch with subscribe/unsubscribe/query protocol handling
709#[allow(clippy::too_many_arguments)] // Reason: WebSocket handler needs connection state, protocol codec, and tenant context
710async fn handle_client_message(
711    text: &str,
712    connection_id: &str,
713    state: &SubscriptionState,
714    codec: &ProtocolCodec,
715    active_operations: &mut HashMap<String, SubscriptionId>,
716    remote_msg_tx: tokio::sync::mpsc::Sender<ServerMessage>,
717    sender: &mut futures::stream::SplitSink<WebSocket, Message>,
718    tenant_id: Option<&str>,
719    principal: Option<&SecurityContext>,
720) -> Result<(), CloseCode> {
721    // remote_msg_tx is only consumed inside the #[cfg(feature = "federation")] block below.
722    // When the feature is disabled the parameter goes unused; suppress the warning.
723    #[cfg(not(feature = "federation"))]
724    let _ = &remote_msg_tx;
725
726    let client_msg: ClientMessage = codec.decode(text).map_err(|e| {
727        warn!(error = %e, "Failed to parse client message");
728        CloseCode::ProtocolError
729    })?;
730
731    match client_msg.parsed_type() {
732        Some(ClientMessageType::Ping) => {
733            let pong = ServerMessage::pong(client_msg.payload);
734            // Best-effort: if the connection is already dead the pong will fail.
735            let _ = send_server_message(codec, sender, pong).await;
736        },
737
738        Some(ClientMessageType::Pong) => {
739            debug!(connection_id = %connection_id, "Received pong");
740        },
741
742        Some(ClientMessageType::Subscribe) => {
743            let payload: SubscribePayload = client_msg.subscription_payload().ok_or_else(|| {
744                warn!("Invalid subscribe payload");
745                CloseCode::ProtocolError
746            })?;
747
748            let op_id = client_msg.id.ok_or_else(|| {
749                warn!("Subscribe message missing operation ID");
750                CloseCode::ProtocolError
751            })?;
752
753            // Check for duplicate operation ID
754            if active_operations.contains_key(&op_id) {
755                warn!(operation_id = %op_id, "Duplicate operation ID");
756                return Err(CloseCode::SubscriberAlreadyExists);
757            }
758
759            // Enforce per-connection subscription limit
760            if let Some(max) = state.max_subscriptions_per_connection {
761                if active_operations.len() >= max as usize {
762                    warn!(
763                        connection_id = %connection_id,
764                        active = active_operations.len(),
765                        max = max,
766                        "Subscription limit reached"
767                    );
768                    WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
769                    let error = ServerMessage::error(
770                        &op_id,
771                        vec![GraphQLError::with_code(
772                            format!("Maximum subscriptions per connection ({max}) reached"),
773                            "SUBSCRIPTION_LIMIT_REACHED",
774                        )],
775                    );
776                    if let Err(e) = send_server_message(codec, sender, error).await {
777                        debug!(connection_id = %connection_id, error = %e, "Could not send subscription limit error to client");
778                    }
779                    return Ok(());
780                }
781            }
782
783            // Extract subscription name from query
784            let Some(subscription_name) = extract_subscription_name(&payload.query) else {
785                let error = ServerMessage::error(
786                    &op_id,
787                    vec![GraphQLError::with_code(
788                        "Could not parse subscription query",
789                        "PARSE_ERROR",
790                    )],
791                );
792                if let Err(e) = send_server_message(codec, sender, error).await {
793                    debug!(connection_id = %connection_id, error = %e, "Could not send parse error to client");
794                }
795                return Ok(());
796            };
797
798            // Call lifecycle on_subscribe hook
799            // HashMap<String, Value> serialization is infallible; the error path cannot occur.
800            let variables_value = serde_json::to_value(&payload.variables)
801                .expect("HashMap<String, serde_json::Value> serialization is infallible");
802
803            // #422: operation-level authorization at subscription establishment.
804            // The per-event delivery does not route through the executor, so the
805            // subscription is authorized once, here, with the connection's principal
806            // (or `None` when anonymous). Fail-closed: a `Deny` or any policy error
807            // rejects the subscription with a `FORBIDDEN` GraphQL-WS error.
808            if let Some(authorizer) = state.authorizer.as_ref() {
809                let ops = [(OperationKind::Subscription, subscription_name.clone())];
810                if let Err(err) =
811                    enforce_authz(authorizer.as_ref(), principal, &ops, Some(&variables_value))
812                {
813                    warn!(
814                        connection_id = %connection_id,
815                        subscription = %subscription_name,
816                        "Operation authorizer denied the subscription"
817                    );
818                    WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
819                    let error = ServerMessage::error(
820                        &op_id,
821                        vec![GraphQLError::with_code(err.to_string(), "FORBIDDEN")],
822                    );
823                    if let Err(e) = send_server_message(codec, sender, error).await {
824                        debug!(connection_id = %connection_id, error = %e, "Could not send authorization denial to client");
825                    }
826                    return Ok(());
827                }
828            }
829
830            if let Err(reason) = state
831                .lifecycle
832                .on_subscribe(&subscription_name, &variables_value, connection_id)
833                .await
834            {
835                warn!(
836                    connection_id = %connection_id,
837                    subscription = %subscription_name,
838                    reason = %reason,
839                    "Lifecycle on_subscribe rejected subscription"
840                );
841                WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
842                let error = ServerMessage::error(
843                    &op_id,
844                    vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REJECTED")],
845                );
846                if let Err(e) = send_server_message(codec, sender, error).await {
847                    debug!(connection_id = %connection_id, error = %e, "Could not send subscription rejection to client");
848                }
849                return Ok(());
850            }
851
852            // Forward to remote subgraph when the subscription field is owned remotely.
853            #[cfg(feature = "federation")]
854            if let Some(subgraph_url) = state.remote_subscription_fields.get(&subscription_name) {
855                use fraiseql_federation::subscription_forwarder::{
856                    ForwardedEvent, SubscriptionForwarder,
857                };
858
859                match SubscriptionForwarder::new(subgraph_url) {
860                    Ok(forwarder) => {
861                        // Create a channel so the forwarder task can send us raw events.
862                        let (event_tx, mut event_rx) =
863                            tokio::sync::mpsc::channel::<ForwardedEvent>(32);
864
865                        // Task 1: run the WebSocket forwarder (sends ForwardedEvent to event_tx).
866                        let fwd_op = op_id.clone();
867                        let fwd_query = payload.query.clone();
868                        let fwd_vars = variables_value.clone();
869                        tokio::spawn(async move {
870                            if let Err(e) =
871                                forwarder.forward(&fwd_op, &fwd_query, fwd_vars, event_tx).await
872                            {
873                                warn!(error = %e, "Remote subscription forwarder failed");
874                            }
875                        });
876
877                        // Task 2: relay ForwardedEvent → ServerMessage → client.
878                        let relay_op = op_id.clone();
879                        let relay_tx = remote_msg_tx.clone();
880                        tokio::spawn(async move {
881                            while let Some(event) = event_rx.recv().await {
882                                let server_msg = match event {
883                                    ForwardedEvent::Next(data) => {
884                                        ServerMessage::next(&relay_op, data)
885                                    },
886                                    ForwardedEvent::Error(errors) => {
887                                        let errors_vec = errors.as_array().map_or_else(
888                                            || {
889                                                vec![GraphQLError::with_code(
890                                                    errors.to_string(),
891                                                    "REMOTE_ERROR",
892                                                )]
893                                            },
894                                            |arr| {
895                                                arr.iter()
896                                                    .map(|e| {
897                                                        GraphQLError::with_code(
898                                                            e.get("message")
899                                                                .and_then(|v| v.as_str())
900                                                                .unwrap_or("Remote subgraph error"),
901                                                            "REMOTE_ERROR",
902                                                        )
903                                                    })
904                                                    .collect()
905                                            },
906                                        );
907                                        ServerMessage::error(&relay_op, errors_vec)
908                                    },
909                                    ForwardedEvent::Complete => ServerMessage::complete(&relay_op),
910                                };
911                                if relay_tx.send(server_msg).await.is_err() {
912                                    break; // Client disconnected
913                                }
914                            }
915                        });
916
917                        WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
918                        info!(
919                            connection_id = %connection_id,
920                            operation_id = %op_id,
921                            subscription = %subscription_name,
922                            "Subscription forwarded to remote subgraph"
923                        );
924                        return Ok(());
925                    },
926                    Err(e) => {
927                        let error = ServerMessage::error(
928                            &op_id,
929                            vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
930                        );
931                        if let Err(send_err) = send_server_message(codec, sender, error).await {
932                            debug!(connection_id = %connection_id, error = %send_err, "Could not send forwarding error to client");
933                        }
934                        return Ok(());
935                    },
936                }
937            }
938
939            // Validate client-provided tenant variable against server-resolved
940            if let Some(server_tid) = tenant_id {
941                if let Some(client_tid) = variables_value.get("tenant_id").and_then(|v| v.as_str())
942                {
943                    if client_tid != server_tid {
944                        let error = ServerMessage::error(
945                            &op_id,
946                            vec![GraphQLError::with_code(
947                                format!(
948                                    "Tenant mismatch: client provided '{client_tid}', server resolved '{server_tid}'"
949                                ),
950                                "TENANT_MISMATCH",
951                            )],
952                        );
953                        if let Err(send_err) = send_server_message(codec, sender, error).await {
954                            debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant mismatch error to client");
955                        }
956                        return Ok(());
957                    }
958                }
959            }
960
961            // M-tenant-ws-suspended: refuse to start a subscription whose tenant
962            // is suspended, mirroring the GraphQL data plane's 503 response.
963            if let (Some(tid), Some(src)) = (tenant_id, state.tenant_status_source.as_ref()) {
964                if src.is_suspended(tid) {
965                    WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
966                    let error = ServerMessage::error(
967                        &op_id,
968                        vec![GraphQLError::with_code(
969                            format!("Tenant '{tid}' is suspended"),
970                            "TENANT_SUSPENDED",
971                        )],
972                    );
973                    if let Err(send_err) = send_server_message(codec, sender, error).await {
974                        debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant-suspended error to client");
975                    }
976                    return Ok(());
977                }
978            }
979
980            // #596: derive the server-owned row-visibility condition for this
981            // subscription from the target entity's `subscription_policy`. Fail-closed —
982            // a policy that applies but whose identity is unresolvable refuses the
983            // subscription rather than falling back to delivering every row.
984            let rls_conditions = match resolve_subscription_rls(
985                state,
986                &subscription_name,
987                principal,
988            )
989            .await
990            {
991                Ok(conditions) => conditions,
992                Err(reason) => {
993                    WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
994                    warn!(
995                        connection_id = %connection_id,
996                        subscription = %subscription_name,
997                        "Row-visibility policy refused the subscription (fail-closed)"
998                    );
999                    let error = ServerMessage::error(
1000                        &op_id,
1001                        vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REFUSED")],
1002                    );
1003                    if let Err(send_err) = send_server_message(codec, sender, error).await {
1004                        debug!(connection_id = %connection_id, error = %send_err, "Could not send row-visibility refusal to client");
1005                    }
1006                    return Ok(());
1007                },
1008            };
1009
1010            // Build context with server-resolved tenant_id
1011            let mut context = serde_json::json!({});
1012            if let Some(tid) = tenant_id {
1013                context["tenant_id"] = serde_json::Value::String(tid.to_string());
1014            }
1015
1016            // Subscribe locally (field is owned by this subgraph). The server-owned
1017            // `rls_conditions` (#596) are enforced on every delivered event (AND
1018            // semantics) and cannot be overridden by client-supplied variables/filters.
1019            match state.manager.subscribe_with_rls(
1020                &subscription_name,
1021                context,
1022                variables_value,
1023                connection_id,
1024                rls_conditions,
1025            ) {
1026                Ok(sub_id) => {
1027                    active_operations.insert(op_id.clone(), sub_id);
1028                    WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
1029                    info!(
1030                        connection_id = %connection_id,
1031                        operation_id = %op_id,
1032                        subscription = %subscription_name,
1033                        "Subscription started"
1034                    );
1035                },
1036                Err(e) => {
1037                    let error = ServerMessage::error(
1038                        &op_id,
1039                        vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
1040                    );
1041                    if let Err(send_err) = send_server_message(codec, sender, error).await {
1042                        debug!(connection_id = %connection_id, error = %send_err, "Could not send subscription error to client");
1043                    }
1044                },
1045            }
1046        },
1047
1048        Some(ClientMessageType::Complete) => {
1049            let op_id = client_msg.id.ok_or_else(|| {
1050                warn!("Complete message missing operation ID");
1051                CloseCode::ProtocolError
1052            })?;
1053
1054            if let Some(sub_id) = active_operations.remove(&op_id) {
1055                if let Err(e) = state.manager.unsubscribe(sub_id) {
1056                    warn!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe; subscription may be leaked");
1057                }
1058                state.lifecycle.on_unsubscribe(&op_id, connection_id).await;
1059                info!(
1060                    connection_id = %connection_id,
1061                    operation_id = %op_id,
1062                    "Subscription completed"
1063                );
1064            }
1065        },
1066
1067        Some(ClientMessageType::ConnectionInit) => {
1068            warn!(connection_id = %connection_id, "Duplicate connection_init");
1069            return Err(CloseCode::TooManyInitRequests);
1070        },
1071
1072        None => {
1073            warn!(message_type = %client_msg.message_type, "Unknown message type");
1074        },
1075        // Reason: non_exhaustive requires catch-all for cross-crate matches
1076        _ => {
1077            warn!(message_type = %client_msg.message_type, "Unrecognized message type");
1078        },
1079    }
1080
1081    Ok(())
1082}
1083
1084/// Send a server message through the codec, handling protocol translation.
1085async fn send_server_message(
1086    codec: &ProtocolCodec,
1087    sender: &mut futures::stream::SplitSink<WebSocket, Message>,
1088    msg: ServerMessage,
1089) -> Result<(), String> {
1090    match codec.encode(&msg) {
1091        Ok(Some(json)) => sender.send(Message::Text(json.into())).await.map_err(|e| e.to_string()),
1092        Ok(None) => Ok(()), // Message suppressed by codec (e.g. pong in legacy mode)
1093        Err(e) => Err(e.to_string()),
1094    }
1095}
1096
1097/// Create a "next" message for a subscription event.
1098///
1099/// When the event carries a Change-Spine envelope (#425), it rides in the
1100/// graphql-transport-ws `ExecutionResult` `extensions` slot as `changeSpine` —
1101/// the spec-blessed, client-ignorable channel — leaving the resolved entity
1102/// `data` untouched. Events without an envelope produce the plain `next` message.
1103fn create_next_message(operation_id: &str, payload: &SubscriptionPayload) -> ServerMessage {
1104    let data = serde_json::json!({
1105        payload.subscription_name.clone(): payload.data
1106    });
1107    match &payload.event.change_spine {
1108        Some(envelope) => {
1109            let extensions = serde_json::json!({ "changeSpine": envelope });
1110            ServerMessage::next_with_extensions(operation_id, data, extensions)
1111        },
1112        None => ServerMessage::next(operation_id, data),
1113    }
1114}
1115
1116/// Extract subscription name from a GraphQL subscription query.
1117pub(crate) fn extract_subscription_name(query: &str) -> Option<String> {
1118    let query = query.trim();
1119
1120    let sub_idx = query.find("subscription")?;
1121    let after_sub = &query[sub_idx + "subscription".len()..];
1122
1123    let brace_idx = after_sub.find('{')?;
1124    let after_brace = after_sub[brace_idx + 1..].trim_start();
1125
1126    let name_end = after_brace
1127        .find(|c: char| !c.is_alphanumeric() && c != '_')
1128        .unwrap_or(after_brace.len());
1129
1130    if name_end == 0 {
1131        return None;
1132    }
1133
1134    Some(after_brace[..name_end].to_string())
1135}
1136
1137#[cfg(test)]
1138mod tests;