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(¶ms, &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;