use std::{
collections::HashMap,
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
time::Duration,
};
use axum::{
extract::{
State,
ws::{Message, WebSocket, WebSocketUpgrade},
},
http::HeaderMap,
response::IntoResponse,
};
use fraiseql_core::{
graphql::ParsedQuery,
runtime::{
SubscriptionId, SubscriptionManager, SubscriptionPayload,
protocol::{
ClientMessage, ClientMessageType, CloseCode, GraphQLError, ServerMessage,
SubscribePayload,
},
},
schema::{CompiledSchema, OwnerCondition, SubscriptionPolicy},
security::{
Authorizer, AuthzOperation, OperationKind, SecurityContext, authorizer::enforce_authz,
},
};
use futures::{SinkExt, StreamExt};
use tokio::sync::broadcast;
use tracing::{debug, error, info, warn};
use crate::middleware::stream_auth::StreamAuthGuard;
use crate::{
extractors::OptionalSecurityContext,
routes::graphql::{DomainRegistry, TenantStatusSource},
subscriptions::{
lifecycle::SubscriptionLifecycle,
protocol::{ProtocolCodec, WsProtocol},
},
};
static WS_CONNECTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
static WS_CONNECTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
static WS_SUBSCRIPTIONS_ACCEPTED: AtomicU64 = AtomicU64::new(0);
static WS_SUBSCRIPTIONS_REJECTED: AtomicU64 = AtomicU64::new(0);
#[must_use]
pub fn subscription_metrics() -> SubscriptionMetrics {
SubscriptionMetrics {
connections_accepted: WS_CONNECTIONS_ACCEPTED.load(Ordering::Relaxed),
connections_rejected: WS_CONNECTIONS_REJECTED.load(Ordering::Relaxed),
subscriptions_accepted: WS_SUBSCRIPTIONS_ACCEPTED.load(Ordering::Relaxed),
subscriptions_rejected: WS_SUBSCRIPTIONS_REJECTED.load(Ordering::Relaxed),
}
}
#[cfg(test)]
pub fn reset_metrics_for_test() {
WS_CONNECTIONS_ACCEPTED.store(0, Ordering::SeqCst);
WS_CONNECTIONS_REJECTED.store(0, Ordering::SeqCst);
WS_SUBSCRIPTIONS_ACCEPTED.store(0, Ordering::SeqCst);
WS_SUBSCRIPTIONS_REJECTED.store(0, Ordering::SeqCst);
}
pub struct SubscriptionMetrics {
pub connections_accepted: u64,
pub connections_rejected: u64,
pub subscriptions_accepted: u64,
pub subscriptions_rejected: u64,
}
const CONNECTION_INIT_TIMEOUT: Duration = Duration::from_secs(5);
const PING_INTERVAL: Duration = Duration::from_secs(30);
pub type LiveSubscriptionPolicies =
Arc<dyn Fn() -> Arc<HashMap<String, SubscriptionPolicy>> + Send + Sync>;
pub type LiveSchema = Arc<dyn Fn() -> Arc<CompiledSchema> + Send + Sync>;
pub type LiveExecutor = Arc<dyn Fn() -> Arc<fraiseql_core::runtime::Executor> + Send + Sync>;
#[derive(Clone)]
pub struct SubscriptionState {
pub manager: Arc<SubscriptionManager>,
pub lifecycle: Arc<dyn SubscriptionLifecycle>,
pub max_subscriptions_per_connection: Option<u32>,
pub domain_registry: Option<Arc<DomainRegistry>>,
pub strict_tenant_validation: bool,
pub authorizer: Option<Arc<dyn Authorizer>>,
pub tenant_status_source: Option<Arc<dyn TenantStatusSource>>,
pub subscription_policies: Arc<HashMap<String, SubscriptionPolicy>>,
pub live_subscription_policies: Option<LiveSubscriptionPolicies>,
pub live_schema: Option<LiveSchema>,
pub live_executor: Option<LiveExecutor>,
#[cfg(feature = "auth")]
pub identity_resolver: Option<Arc<crate::identity::IdentityResolver>>,
pub service_account_authenticator:
Option<Arc<crate::service_account::ServiceAccountAuthenticator>>,
pub policy_reload: Option<tokio::sync::watch::Receiver<u64>>,
pub drain: Option<tokio::sync::watch::Receiver<bool>>,
pub revocation_manager: Option<Arc<crate::token_revocation::TokenRevocationManager>>,
pub auth_recheck_interval: Duration,
}
impl SubscriptionState {
pub fn new(manager: Arc<SubscriptionManager>) -> Self {
Self {
manager,
lifecycle: Arc::new(crate::subscriptions::lifecycle::NoopLifecycle),
max_subscriptions_per_connection: None,
domain_registry: None,
strict_tenant_validation: false,
authorizer: None,
tenant_status_source: None,
subscription_policies: Arc::new(HashMap::new()),
live_subscription_policies: None,
live_schema: None,
live_executor: None,
#[cfg(feature = "auth")]
identity_resolver: None,
service_account_authenticator: None,
policy_reload: None,
drain: None,
revocation_manager: None,
auth_recheck_interval: Duration::from_secs(
crate::server_config::defaults::default_subscription_auth_recheck_secs(),
),
}
}
#[must_use]
pub fn with_policy_reload(
mut self,
policy_reload: Option<tokio::sync::watch::Receiver<u64>>,
) -> Self {
self.policy_reload = policy_reload;
self
}
#[must_use]
pub fn with_drain_signal(mut self, drain: Option<tokio::sync::watch::Receiver<bool>>) -> Self {
self.drain = drain;
self
}
#[must_use]
pub fn with_revocation_manager(
mut self,
manager: Option<Arc<crate::token_revocation::TokenRevocationManager>>,
) -> Self {
self.revocation_manager = manager;
self
}
#[must_use]
pub const fn with_auth_recheck_interval(mut self, interval: Duration) -> Self {
self.auth_recheck_interval = interval;
self
}
#[must_use]
pub fn with_subscription_policies(
mut self,
policies: Arc<HashMap<String, SubscriptionPolicy>>,
) -> Self {
self.subscription_policies = policies;
self
}
#[must_use]
pub fn with_live_subscription_policies(
mut self,
live: Option<LiveSubscriptionPolicies>,
) -> Self {
self.live_subscription_policies = live;
self
}
#[must_use]
pub fn with_live_schema(mut self, live: Option<LiveSchema>) -> Self {
self.live_schema = live;
self
}
#[must_use]
pub fn with_live_executor(mut self, live: Option<LiveExecutor>) -> Self {
self.live_executor = live;
self
}
#[must_use]
pub fn with_service_account_authenticator(
mut self,
authenticator: Option<Arc<crate::service_account::ServiceAccountAuthenticator>>,
) -> Self {
self.service_account_authenticator = authenticator;
self
}
#[cfg(feature = "auth")]
#[must_use]
pub fn with_identity_resolver(
mut self,
resolver: Option<Arc<crate::identity::IdentityResolver>>,
) -> Self {
self.identity_resolver = resolver;
self
}
#[must_use]
pub fn with_tenant_status_source(
mut self,
source: Option<Arc<dyn TenantStatusSource>>,
) -> Self {
self.tenant_status_source = source;
self
}
#[must_use]
pub fn with_authorizer(mut self, authorizer: Option<Arc<dyn Authorizer>>) -> Self {
self.authorizer = authorizer;
self
}
#[must_use]
pub fn with_tenant_context(
mut self,
domain_registry: Arc<DomainRegistry>,
strict_tenant_validation: bool,
) -> Self {
self.domain_registry = Some(domain_registry);
self.strict_tenant_validation = strict_tenant_validation;
self
}
#[must_use]
pub fn with_lifecycle(mut self, lifecycle: Arc<dyn SubscriptionLifecycle>) -> Self {
self.lifecycle = lifecycle;
self
}
#[must_use]
pub const fn with_max_subscriptions(mut self, max: Option<u32>) -> Self {
self.max_subscriptions_per_connection = max;
self
}
}
#[must_use]
pub fn build_subscription_policies(schema: &CompiledSchema) -> HashMap<String, SubscriptionPolicy> {
let mut policies = HashMap::new();
for sub in &schema.subscriptions {
if let Some(type_def) = schema.types.iter().find(|t| t.name.as_str() == sub.return_type) {
if let Some(policy) = &type_def.subscription_policy {
policies.insert(sub.name.clone(), policy.clone());
}
}
}
policies
}
async fn resolve_subscription_rls(
state: &SubscriptionState,
subscription_name: &str,
principal: Option<&SecurityContext>,
) -> Result<Vec<(String, serde_json::Value)>, String> {
let policies = state
.live_subscription_policies
.as_ref()
.map_or_else(|| Arc::clone(&state.subscription_policies), |live| live());
let Some(policy) = policies.get(subscription_name) else {
return Ok(Vec::new());
};
#[cfg(feature = "auth")]
let enriched = enrich_principal(state, principal).await;
#[cfg(feature = "auth")]
let effective = enriched.as_ref().or(principal);
#[cfg(not(feature = "auth"))]
let effective = principal;
derive_policy_conditions(policy, effective)
}
#[cfg(feature = "auth")]
async fn enrich_principal(
state: &SubscriptionState,
principal: Option<&SecurityContext>,
) -> Option<SecurityContext> {
let resolver = state.identity_resolver.as_ref()?;
let mut ctx = principal?.clone();
let _ = crate::identity::resolve_request_identity(Some(resolver), Some(&mut ctx)).await;
Some(ctx)
}
fn derive_policy_conditions(
policy: &SubscriptionPolicy,
principal: Option<&SecurityContext>,
) -> Result<Vec<(String, serde_json::Value)>, String> {
let empty = HashMap::new();
let attributes = principal.map_or(&empty, |ctx| &ctx.attributes);
let roles: &[String] = principal.map_or(&[], |ctx| ctx.roles.as_slice());
match policy.derive(attributes, roles) {
OwnerCondition::Bypass => Ok(Vec::new()),
OwnerCondition::Eq { field, value } => Ok(vec![(field, value)]),
OwnerCondition::Refuse(reason) => Err(reason),
}
}
pub async fn subscription_handler(
headers: HeaderMap,
OptionalSecurityContext(security_context): OptionalSecurityContext,
token_claims: Option<axum::Extension<crate::middleware::oidc_auth::SessionTokenClaims>>,
ws: WebSocketUpgrade,
State(state): State<SubscriptionState>,
) -> impl IntoResponse {
let protocol_header = headers.get("sec-websocket-protocol").and_then(|v| v.to_str().ok());
let protocol = match protocol_header {
None => WsProtocol::GraphqlTransportWs,
Some(header) => {
if let Some(p) = WsProtocol::from_header(Some(header)) {
p
} else {
warn!(header = %header, "Unknown WebSocket sub-protocol requested");
return axum::http::StatusCode::BAD_REQUEST.into_response();
}
},
};
let mut security_context = security_context;
if let Some(sa_auth) = state.service_account_authenticator.as_ref() {
match sa_auth.resolve(&headers, security_context.is_some()) {
crate::service_account::SaAuth::NoSecret => {},
crate::service_account::SaAuth::Authenticated(ctx) => security_context = Some(*ctx),
crate::service_account::SaAuth::Ambiguous
| crate::service_account::SaAuth::Unmatched => {
warn!(
"Subscription upgrade rejected: ambiguous or unmatched service-account secret"
);
return axum::http::StatusCode::UNAUTHORIZED.into_response();
},
}
}
let tenant_id = match resolve_subscription_tenant(security_context.as_ref(), &headers, &state) {
Ok(tenant_id) => tenant_id,
Err(e) => {
warn!(error = %e, "Subscription tenant resolution rejected the upgrade");
return axum::http::StatusCode::BAD_REQUEST.into_response();
},
};
ws.protocols([protocol.as_str()])
.on_upgrade(move |socket| {
handle_subscription_connection(
socket,
state,
protocol,
tenant_id,
security_context,
token_claims.map(|axum::Extension(claims)| claims),
)
})
.into_response()
}
fn resolve_subscription_tenant(
security_context: Option<&SecurityContext>,
headers: &HeaderMap,
state: &SubscriptionState,
) -> fraiseql_error::Result<Option<String>> {
super::graphql::TenantKeyResolver::resolve(
security_context,
headers,
state.domain_registry.as_deref(),
state.strict_tenant_validation,
)
}
enum InitWait {
Init(Box<ClientMessage>),
ClientClosed,
Violation(CloseCode),
}
#[allow(clippy::cognitive_complexity)] async fn handle_subscription_connection(
socket: WebSocket,
state: SubscriptionState,
protocol: WsProtocol,
tenant_id: Option<String>,
principal: Option<SecurityContext>,
token_claims: Option<crate::middleware::oidc_auth::SessionTokenClaims>,
) {
let connection_id = uuid::Uuid::new_v4().to_string();
let codec = ProtocolCodec::new(protocol);
info!(
connection_id = %connection_id,
protocol = %protocol.as_str(),
"WebSocket connection established"
);
let (mut sender, mut receiver) = socket.split();
let init_result = tokio::time::timeout(CONNECTION_INIT_TIMEOUT, async {
while let Some(msg) = receiver.next().await {
match msg {
Ok(Message::Text(text)) => match codec.decode(&text) {
Ok(client_msg) => {
if client_msg.parsed_type() == Some(ClientMessageType::ConnectionInit) {
return InitWait::Init(Box::new(client_msg));
}
warn!(
message_type = %client_msg.message_type,
"Message before connection_init; closing 4401"
);
return InitWait::Violation(CloseCode::Unauthorized);
},
Err(e) => {
warn!(error = %e, "Undecodable message before connection_init; closing 4400");
return InitWait::Violation(CloseCode::BadRequest);
},
},
Ok(Message::Close(_)) => return InitWait::ClientClosed,
Err(e) => {
error!(error = %e, "WebSocket error during init");
return InitWait::ClientClosed;
},
_ => {},
}
}
InitWait::ClientClosed
})
.await;
let _init_payload = match init_result {
Ok(InitWait::Init(msg)) => {
let params = msg.payload.clone().unwrap_or(serde_json::json!({}));
if let Err(reason) = state.lifecycle.on_connect(¶ms, &connection_id).await {
warn!(
connection_id = %connection_id,
reason = %reason,
"Lifecycle on_connect rejected connection"
);
WS_CONNECTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
let _ = sender
.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: 4400,
reason: reason.into(),
})))
.await;
return;
}
let ack = ServerMessage::connection_ack(None);
if let Err(send_err) = send_server_message(&codec, &mut sender, ack).await {
error!(connection_id = %connection_id, error = %send_err, "Failed to send connection_ack");
return;
}
WS_CONNECTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
info!(connection_id = %connection_id, "Connection initialized");
msg.payload
},
Ok(InitWait::ClientClosed) => {
warn!(connection_id = %connection_id, "Connection closed during init");
return;
},
Ok(InitWait::Violation(code)) => {
WS_CONNECTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
let _ = sender
.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: code.code(),
reason: code.reason().into(),
})))
.await;
return;
},
Err(_) => {
warn!(connection_id = %connection_id, "Connection init timeout");
let _ = sender
.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: CloseCode::ConnectionInitTimeout.code(),
reason: CloseCode::ConnectionInitTimeout.reason().into(),
})))
.await;
return;
},
};
let mut active_operations: HashMap<String, ActiveOperation> = HashMap::new();
let mut policy_reload = state.policy_reload.clone();
let mut drain = state.drain.clone();
let mut event_receiver = state.manager.receiver();
let mut ping_interval = tokio::time::interval(PING_INTERVAL);
ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
let auth_guard =
StreamAuthGuard::new(principal.as_ref(), token_claims, state.revocation_manager.clone());
let auth_recheck_enabled = auth_guard.applies() && !state.auth_recheck_interval.is_zero();
let mut auth_recheck =
tokio::time::interval(state.auth_recheck_interval.max(Duration::from_millis(1)));
auth_recheck.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
auth_recheck.tick().await;
loop {
tokio::select! {
msg = receiver.next() => {
match msg {
Some(Ok(Message::Text(text))) => {
if let Err(close_code) = handle_client_message(
&text,
&connection_id,
&state,
&codec,
&mut active_operations,
&mut sender,
tenant_id.as_deref(),
principal.as_ref(),
).await {
let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: close_code.code(),
reason: close_code.reason().into(),
}))).await;
break;
}
}
Some(Ok(Message::Ping(data))) => {
let _ = sender.send(Message::Pong(data)).await;
}
Some(Ok(Message::Close(_))) => {
info!(connection_id = %connection_id, "Client closed connection");
break;
}
Some(Err(e)) => {
error!(connection_id = %connection_id, error = %e, "WebSocket error");
break;
}
None => {
info!(connection_id = %connection_id, "WebSocket stream ended");
break;
}
_ => {}
}
}
draining = drain_signalled(&mut drain), if drain.is_some() => {
if draining {
info!(
connection_id = %connection_id,
active_operations = active_operations.len(),
"Server draining; completing active subscriptions"
);
for (op_id, op) in active_operations.drain() {
let complete = ServerMessage::complete(&op_id);
if let Err(e) = send_server_message(&codec, &mut sender, complete).await {
debug!(connection_id = %connection_id, error = %e, "Could not send Complete during drain");
}
if let Err(e) = state.manager.unsubscribe(op.subscription_id) {
debug!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe during drain");
}
state.lifecycle.on_unsubscribe(&op_id, &connection_id).await;
}
let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: CloseCode::GoingAway.code(),
reason: CloseCode::GoingAway.reason().into(),
}))).await;
break;
}
drain = None;
}
changed = policy_watch_changed(&mut policy_reload), if policy_reload.is_some() => {
if changed {
rederive_operations_after_policy_reload(
&state,
&codec,
&mut sender,
&mut active_operations,
&connection_id,
principal.as_ref(),
).await;
} else {
policy_reload = None;
}
}
_ = auth_recheck.tick(), if auth_recheck_enabled => {
if let Err(reason) = auth_guard.check().await {
warn!(
connection_id = %connection_id,
reason,
"Mid-stream authorization re-check failed; closing WebSocket"
);
let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: CloseCode::Unauthorized.code(),
reason: reason.into(),
}))).await;
break;
}
}
event = event_receiver.recv() => {
match event {
Ok(payload) => {
if auth_guard.expired() {
warn!(
connection_id = %connection_id,
"Token expired at delivery time; closing WebSocket"
);
let _ = sender.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: CloseCode::Unauthorized.code(),
reason: "Token expired".into(),
}))).await;
break;
}
let tenant_matches = match (
tenant_id.as_deref(),
payload.event.tenant_id.as_deref(),
) {
(Some(conn_tid), Some(evt_tid)) => conn_tid == evt_tid,
_ => true, };
let tenant_active = match (
tenant_id.as_deref(),
state.tenant_status_source.as_ref(),
) {
(Some(tid), Some(src)) => !src.is_suspended(tid),
_ => true,
};
if tenant_matches && tenant_active {
if let Some((op_id, op)) = active_operations
.iter()
.find(|(_, op)| op.subscription_id == payload.subscription_id)
{
let msg = create_next_message(op_id, &op.response_key, &payload);
if send_server_message(&codec, &mut sender, msg).await.is_err() {
warn!(connection_id = %connection_id, "Failed to send event");
break;
}
}
}
}
Err(broadcast::error::RecvError::Lagged(n)) => {
warn!(
connection_id = %connection_id,
lagged = n,
active_operations = active_operations.len(),
"Event receiver lagged; terminating operations with EVENTS_LAGGED"
);
terminate_operations_after_lag(
&state,
&codec,
&mut sender,
&mut active_operations,
&connection_id,
n,
)
.await;
}
Err(broadcast::error::RecvError::Closed) => {
error!(connection_id = %connection_id, "Event channel closed");
break;
}
}
}
_ = ping_interval.tick() => {
let msg = ServerMessage::ping(None);
if send_server_message(&codec, &mut sender, msg).await.is_err() {
warn!(connection_id = %connection_id, "Failed to send ping/keepalive");
break;
}
}
}
}
state.manager.unsubscribe_connection(&connection_id);
state.lifecycle.on_disconnect(&connection_id).await;
info!(connection_id = %connection_id, "WebSocket connection closed");
}
struct ActiveOperation {
subscription_id: SubscriptionId,
subscription_name: String,
response_key: String,
planned_from: Option<(ParsedQuery, serde_json::Value)>,
}
async fn policy_watch_changed(rx: &mut Option<tokio::sync::watch::Receiver<u64>>) -> bool {
match rx {
Some(rx) => rx.changed().await.is_ok(),
None => std::future::pending().await,
}
}
async fn drain_signalled(rx: &mut Option<tokio::sync::watch::Receiver<bool>>) -> bool {
match rx {
Some(rx) => {
rx.wait_for(|draining| *draining).await.is_ok()
},
None => std::future::pending().await,
}
}
async fn rederive_operations_after_policy_reload(
state: &SubscriptionState,
codec: &ProtocolCodec,
sender: &mut futures::stream::SplitSink<WebSocket, Message>,
active_operations: &mut HashMap<String, ActiveOperation>,
connection_id: &str,
principal: Option<&SecurityContext>,
) {
let ops: Vec<(String, SubscriptionId, String, Option<(ParsedQuery, serde_json::Value)>)> =
active_operations
.iter()
.map(|(op_id, op)| {
(
op_id.clone(),
op.subscription_id,
op.subscription_name.clone(),
op.planned_from.clone(),
)
})
.collect();
for (op_id, sub_id, name, planned_from) in ops {
let replanned = match (planned_from, state.live_executor.as_ref()) {
(Some((document, variables)), Some(live)) => {
match live().plan_subscription(&document, Some(&variables), principal) {
Ok(plan) => Some(Ok(Arc::new(plan))),
Err(refusal) => Some(Err(refusal.to_string())),
}
},
_ => None,
};
let derived = match replanned {
Some(Err(reason)) => Err(reason),
Some(Ok(plan)) => resolve_subscription_rls(state, &name, principal)
.await
.map(|conditions| (conditions, Some(plan))),
None => resolve_subscription_rls(state, &name, principal)
.await
.map(|conditions| (conditions, None)),
};
match derived.and_then(|(conditions, plan)| {
if let Some(plan) = plan {
state.manager.replace_plan(sub_id, plan).map_err(|e| e.to_string())?;
}
Ok(conditions)
}) {
Ok(conditions) => {
if let Err(e) = state.manager.update_rls_conditions(sub_id, conditions) {
debug!(
connection_id = %connection_id,
operation_id = %op_id,
error = %e,
"Subscription gone during policy re-derivation"
);
active_operations.remove(&op_id);
}
},
Err(reason) => {
warn!(
connection_id = %connection_id,
operation_id = %op_id,
subscription = %name,
"Hot-reloaded row-visibility policy refused an active subscription \
(fail-closed); terminating the operation"
);
active_operations.remove(&op_id);
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REFUSED")],
);
if let Err(e) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %e, "Could not send policy refusal to client");
}
if let Err(e) = state.manager.unsubscribe(sub_id) {
debug!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe policy-refused operation");
}
state.lifecycle.on_unsubscribe(&op_id, connection_id).await;
},
}
}
}
#[allow(clippy::cognitive_complexity)] #[allow(clippy::too_many_arguments)] async fn handle_client_message(
text: &str,
connection_id: &str,
state: &SubscriptionState,
codec: &ProtocolCodec,
active_operations: &mut HashMap<String, ActiveOperation>,
sender: &mut futures::stream::SplitSink<WebSocket, Message>,
tenant_id: Option<&str>,
principal: Option<&SecurityContext>,
) -> Result<(), CloseCode> {
let client_msg: ClientMessage = codec.decode(text).map_err(|e| {
warn!(error = %e, "Failed to parse client message");
CloseCode::BadRequest
})?;
match client_msg.parsed_type() {
Some(ClientMessageType::Ping) => {
let pong = ServerMessage::pong(client_msg.payload);
let _ = send_server_message(codec, sender, pong).await;
},
Some(ClientMessageType::Pong) => {
debug!(connection_id = %connection_id, "Received pong");
},
Some(ClientMessageType::Subscribe) => {
let payload: SubscribePayload = client_msg.subscription_payload().ok_or_else(|| {
warn!("Invalid subscribe payload");
CloseCode::BadRequest
})?;
let op_id = client_msg.id.ok_or_else(|| {
warn!("Subscribe message missing operation ID");
CloseCode::BadRequest
})?;
if active_operations.contains_key(&op_id) {
warn!(operation_id = %op_id, "Duplicate operation ID");
return Err(CloseCode::SubscriberAlreadyExists);
}
if let Some(max) = state.max_subscriptions_per_connection {
if active_operations.len() >= max as usize {
warn!(
connection_id = %connection_id,
active = active_operations.len(),
max = max,
"Subscription limit reached"
);
WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(
format!("Maximum subscriptions per connection ({max}) reached"),
"SUBSCRIPTION_LIMIT_REACHED",
)],
);
if let Err(e) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %e, "Could not send subscription limit error to client");
}
return Ok(());
}
}
let document = extract_subscription_root(&payload.query).and_then(|(root, parsed)| {
let schema = state.live_schema.as_ref().map(|f| f());
validate_subscription_variables(&parsed, schema.as_deref())?;
Ok((root, parsed))
});
let (
SubscriptionRoot {
name: subscription_name,
response_key,
},
parsed,
) = match document {
Ok(root) => root,
Err(refusal) => {
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(
refusal.message().to_string(),
refusal.code(),
)],
);
if let Err(e) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %e, "Could not send document error to client");
}
return Ok(());
},
};
let variables_value = serde_json::to_value(&payload.variables)
.expect("HashMap<String, serde_json::Value> serialization is infallible");
if let Some(authorizer) = state.authorizer.as_ref() {
let live = state.live_schema.as_ref().map(|f| f());
let target = live
.as_deref()
.and_then(|schema| schema.find_subscription(&subscription_name))
.map(|sub| sub.return_type.clone());
let ops = [AuthzOperation::root(
OperationKind::Subscription,
subscription_name.clone(),
target.as_deref(),
)];
if let Err(err) =
enforce_authz(authorizer.as_ref(), principal, &ops, Some(&variables_value))
{
warn!(
connection_id = %connection_id,
subscription = %subscription_name,
"Operation authorizer denied the subscription"
);
WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(err.to_string(), "FORBIDDEN")],
);
if let Err(e) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %e, "Could not send authorization denial to client");
}
return Ok(());
}
}
let plan = match state.live_executor.as_ref() {
Some(live) => {
match live().plan_subscription(&parsed, Some(&variables_value), principal) {
Ok(plan) => Some(Arc::new(plan)),
Err(err) => {
WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
let code = if matches!(
err,
fraiseql_core::error::FraiseQLError::Authorization { .. }
) {
"FORBIDDEN"
} else {
"SUBSCRIPTION_REFUSED"
};
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(err.to_string(), code)],
);
if let Err(e) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %e, "Could not send subscription plan refusal to client");
}
return Ok(());
},
}
},
_ => None,
};
if let Err(reason) = state
.lifecycle
.on_subscribe(&subscription_name, &variables_value, connection_id)
.await
{
warn!(
connection_id = %connection_id,
subscription = %subscription_name,
reason = %reason,
"Lifecycle on_subscribe rejected subscription"
);
WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REJECTED")],
);
if let Err(e) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %e, "Could not send subscription rejection to client");
}
return Ok(());
}
if let Some(server_tid) = tenant_id {
if let Some(client_tid) = variables_value.get("tenant_id").and_then(|v| v.as_str())
{
if client_tid != server_tid {
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(
format!(
"Tenant mismatch: client provided '{client_tid}', server resolved '{server_tid}'"
),
"TENANT_MISMATCH",
)],
);
if let Err(send_err) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant mismatch error to client");
}
return Ok(());
}
}
}
if let (Some(tid), Some(src)) = (tenant_id, state.tenant_status_source.as_ref()) {
if src.is_suspended(tid) {
WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(
format!("Tenant '{tid}' is suspended"),
"TENANT_SUSPENDED",
)],
);
if let Err(send_err) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %send_err, "Could not send tenant-suspended error to client");
}
return Ok(());
}
}
let rls_conditions = match resolve_subscription_rls(
state,
&subscription_name,
principal,
)
.await
{
Ok(conditions) => conditions,
Err(reason) => {
WS_SUBSCRIPTIONS_REJECTED.fetch_add(1, Ordering::Relaxed);
warn!(
connection_id = %connection_id,
subscription = %subscription_name,
"Row-visibility policy refused the subscription (fail-closed)"
);
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(reason, "SUBSCRIPTION_REFUSED")],
);
if let Err(send_err) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %send_err, "Could not send row-visibility refusal to client");
}
return Ok(());
},
};
let mut context = serde_json::json!({});
if let Some(tid) = tenant_id {
context["tenant_id"] = serde_json::Value::String(tid.to_string());
}
let planned_from = plan.as_ref().map(|_| (parsed.clone(), variables_value.clone()));
let registered = match plan {
Some(plan) => state.manager.subscribe_planned(
plan,
context,
variables_value,
connection_id,
rls_conditions,
),
None => state.manager.subscribe_with_rls(
&subscription_name,
context,
variables_value,
connection_id,
rls_conditions,
),
};
match registered {
Ok(sub_id) => {
active_operations.insert(
op_id.clone(),
ActiveOperation {
subscription_id: sub_id,
subscription_name: subscription_name.clone(),
response_key: response_key.clone(),
planned_from,
},
);
WS_SUBSCRIPTIONS_ACCEPTED.fetch_add(1, Ordering::Relaxed);
info!(
connection_id = %connection_id,
operation_id = %op_id,
subscription = %subscription_name,
"Subscription started"
);
},
Err(e) => {
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(e.to_string(), "SUBSCRIPTION_ERROR")],
);
if let Err(send_err) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %send_err, "Could not send subscription error to client");
}
},
}
},
Some(ClientMessageType::Complete) => {
let op_id = client_msg.id.ok_or_else(|| {
warn!("Complete message missing operation ID");
CloseCode::ProtocolError
})?;
if let Some(op) = active_operations.remove(&op_id) {
if let Err(e) = state.manager.unsubscribe(op.subscription_id) {
warn!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe; subscription may be leaked");
}
state.lifecycle.on_unsubscribe(&op_id, connection_id).await;
info!(
connection_id = %connection_id,
operation_id = %op_id,
"Subscription completed"
);
}
},
Some(ClientMessageType::ConnectionInit) => {
warn!(connection_id = %connection_id, "Duplicate connection_init");
return Err(CloseCode::TooManyInitRequests);
},
Some(ClientMessageType::ConnectionTerminate) => {
info!(connection_id = %connection_id, "connection_terminate received; closing");
return Err(CloseCode::Normal);
},
None => {
warn!(message_type = %client_msg.message_type, "Unknown message type");
},
_ => {
warn!(message_type = %client_msg.message_type, "Unrecognized message type");
},
}
Ok(())
}
async fn terminate_operations_after_lag(
state: &SubscriptionState,
codec: &ProtocolCodec,
sender: &mut futures::stream::SplitSink<WebSocket, Message>,
active_operations: &mut HashMap<String, ActiveOperation>,
connection_id: &str,
lagged: u64,
) {
for (op_id, op) in active_operations.drain() {
let error = ServerMessage::error(
&op_id,
vec![GraphQLError::with_code(
format!(
"Event stream lagged: {lagged} events were dropped for this connection. \
Re-subscribe and re-query to resynchronize."
),
"EVENTS_LAGGED",
)],
);
if let Err(e) = send_server_message(codec, sender, error).await {
debug!(connection_id = %connection_id, error = %e, "Could not send EVENTS_LAGGED to client");
}
if let Err(e) = state.manager.unsubscribe(op.subscription_id) {
warn!(connection_id = %connection_id, operation_id = %op_id, error = %e, "Failed to unsubscribe lagged operation");
}
state.lifecycle.on_unsubscribe(&op_id, connection_id).await;
}
}
async fn send_server_message(
codec: &ProtocolCodec,
sender: &mut futures::stream::SplitSink<WebSocket, Message>,
msg: ServerMessage,
) -> Result<(), String> {
match codec.encode(&msg) {
Ok(Some(json)) => sender.send(Message::Text(json.into())).await.map_err(|e| e.to_string()),
Ok(None) => Ok(()), Err(e) => Err(e.to_string()),
}
}
fn create_next_message(
operation_id: &str,
response_key: &str,
payload: &SubscriptionPayload,
) -> ServerMessage {
let data = serde_json::json!({
response_key.to_owned(): payload.data
});
match &payload.event.change_spine {
Some(envelope) => {
let extensions = serde_json::json!({ "changeSpine": envelope });
ServerMessage::next_with_extensions(operation_id, data, extensions)
},
None => ServerMessage::next(operation_id, data),
}
}
pub(crate) struct SubscriptionRoot {
pub name: String,
pub response_key: String,
}
pub(crate) fn extract_subscription_root(
query: &str,
) -> Result<(SubscriptionRoot, ParsedQuery), SubscriptionDocumentError> {
use graphql_parser::query::{Definition, OperationDefinition, Selection};
let parse_err = |msg: &str| {
SubscriptionDocumentError::Parse(format!("Could not parse subscription: {msg}"))
};
let doc = fraiseql_core::graphql::complexity::parse_graphql_document(query)
.map_err(|e| parse_err(&e.to_string()))?;
let mut sub_ops = doc.definitions.iter().filter_map(|def| match def {
Definition::Operation(OperationDefinition::Subscription(sub)) => Some(sub),
_ => None,
});
let sub = sub_ops
.next()
.ok_or_else(|| parse_err("no subscription operation in document"))?;
if sub_ops.next().is_some() {
return Err(parse_err(
"more than one subscription operation; one connection operation serves exactly one",
));
}
let mut root_fields = sub.selection_set.items.iter().filter_map(|sel| match sel {
Selection::Field(field) => Some(field),
_ => None,
});
let first = root_fields.next().ok_or_else(|| parse_err("subscription has no root field"))?;
if root_fields.next().is_some() {
return Err(parse_err(
"more than one root field; one connection operation serves exactly one",
));
}
let root = SubscriptionRoot {
name: first.name.clone(),
response_key: first.alias.clone().unwrap_or_else(|| first.name.clone()),
};
let parsed = fraiseql_core::graphql::parse_selected_operation(
&OperationDefinition::Subscription(sub.clone()),
&doc,
query,
)
.map_err(|e| parse_err(&e.to_string()))?;
Ok((root, parsed))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum SubscriptionDocumentError {
Parse(String),
Validation(String),
}
impl SubscriptionDocumentError {
pub(crate) const fn code(&self) -> &'static str {
match self {
Self::Parse(_) => "PARSE_ERROR",
Self::Validation(_) => "VALIDATION_ERROR",
}
}
pub(crate) fn message(&self) -> &str {
match self {
Self::Parse(m) | Self::Validation(m) => m,
}
}
}
pub(crate) fn validate_subscription_variables(
parsed: &ParsedQuery,
schema: Option<&CompiledSchema>,
) -> Result<(), SubscriptionDocumentError> {
use fraiseql_core::runtime::{
collect_variable_references, validate_variable_types, validate_variable_uses,
validate_variables_used,
};
let fail = |e: fraiseql_core::error::FraiseQLError| {
SubscriptionDocumentError::Validation(e.to_string())
};
let operation_name = parsed.operation_name.as_deref();
let defined: Vec<String> = parsed.variables.iter().map(|v| v.name.clone()).collect();
let referenced = collect_variable_references(parsed).map_err(fail)?;
validate_variable_uses(operation_name, &defined, &referenced).map_err(fail)?;
if let Some(schema) = schema {
validate_variable_types(schema, operation_name, &parsed.variables).map_err(fail)?;
}
validate_variables_used(operation_name, &defined, &referenced).map_err(fail)?;
Ok(())
}
#[cfg(test)]
mod tests;