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