use crate::query_planner::planner::plan_nodes::CustomScalarPaths;
use crate::telemetry::logging::targets;
use bytes::Bytes;
use hyper_rustls::ConfigBuilderExt;
use std::{cell::RefCell, rc::Rc, sync::Arc};
use futures::{stream::LocalBoxStream, StreamExt};
use ntex::{
channel::{mpsc, oneshot},
io::Sealed,
rt,
ws::{self, error::WsError, WsClient as NtexWsClient, WsConnection, WsSink},
SharedCfg,
};
use tracing::{debug, error, trace};
use crate::executor::{
executors::graphql_transport_ws::{SubscribePayload, WS_SUBPROTOCOL},
response::subgraph_response::SubgraphResponse,
};
use crate::executor::{
executors::{
graphql_transport_ws::{ClientMessage, CloseCode, ConnectionInitPayload, ServerMessage},
websocket_common::{
handshake_timeout, heartbeat, parse_frame_to_text, FrameNotParsedToText, WsState,
},
},
response::graphql_error::GraphQLError,
};
#[derive(Debug, thiserror::Error)]
pub enum WsConnectError {
#[error("Missing schema from WebSocket URI: {0}")]
MissingUriSchema(String),
#[error("Wrong WebSocket URI schema: {0}")]
WrongUriSchema(String),
#[error("WebSocket client error: {0}")]
Client(#[from] ws::error::WsClientError),
#[error("WebSocket client builder error: {0}")]
BuilderError(String),
#[error("Failed to load native TLS certificates: {0}")]
NativeTlsCertificatesError(String),
}
#[derive(Debug, thiserror::Error)]
pub enum WsInitError {
#[error("Connection acknowledgement receiver failed")]
ConnectionAckReceiverError,
#[error("Connection acknowledgement receiver closed")]
ConnectionAckReceiverClosed,
#[error("Connection closed before acknowledgement")]
ConnectionClosedBeforeAck,
#[error("Invalid message received during acknowledgement")]
InvalidMessage,
#[error("Wrong message received before connection acknowledgement")]
WrongMessageBeforeAck,
#[error("Failed to send connection initialization message: {0}")]
SendFailed(#[from] ws::error::ProtocolError),
}
#[derive(Clone, Debug, thiserror::Error)]
pub enum WsClientError {
#[error("Connection closed")]
ConnectionClosed,
#[error("Message dispatcher closed")]
MessageDispatcherClosed,
#[error("Failed to deserialize payload")]
FailedToDeserializePayload,
#[error("Failed to send WebSocket message: {0}")]
SendFailed(#[from] ws::error::ProtocolError),
#[error("WebSocket subscription ID exhausted")]
SubscriptionIdExhausted,
}
impl WsClientError {
pub fn error_code(&self) -> &'static str {
match self {
WsClientError::ConnectionClosed => "WS_CONNECTION_CLOSED",
WsClientError::MessageDispatcherClosed => "WS_MESSAGE_DISPATCHER_CLOSED",
WsClientError::FailedToDeserializePayload => "WS_FAILED_TO_DESERIALIZE_PAYLOAD",
WsClientError::SendFailed(_) => "WS_SEND_FAILED",
WsClientError::SubscriptionIdExhausted => "WS_SUBSCRIPTION_ID_EXHAUSTED",
}
}
}
impl From<WsClientError> for SubgraphResponse<'static> {
fn from(err: WsClientError) -> Self {
SubgraphResponse {
errors: Some(vec![GraphQLError::from_message_and_code(
err.to_string(),
err.error_code(),
)]),
..Default::default()
}
}
}
pub async fn connect(
uri: &http::Uri,
custom_tls_config: Option<Arc<rustls::ClientConfig>>,
) -> Result<WsConnection<ntex::io::Sealed>, WsConnectError> {
let scheme = uri
.scheme_str()
.ok_or_else(|| WsConnectError::MissingUriSchema(uri.to_string()))?;
if scheme == "wss" {
let tls_config = match custom_tls_config {
Some(config) => config,
None => Arc::new(
rustls::ClientConfig::builder()
.with_native_roots()
.map_err(|e| WsConnectError::NativeTlsCertificatesError(e.to_string()))?
.with_no_client_auth(),
),
};
let ws_client = NtexWsClient::builder(uri)
.max_frame_size(16 * 1024 * 1024) .protocols([WS_SUBPROTOCOL])
.timeout(ntex::time::Seconds(60))
.rustls(tls_config)
.take()
.build(SharedCfg::default())
.await
.map_err(|e| WsConnectError::BuilderError(e.to_string()))?;
Ok(ws_client.connect().await?.seal())
} else if scheme == "ws" {
let ws_client = NtexWsClient::builder(uri)
.max_frame_size(16 * 1024 * 1024) .protocols([WS_SUBPROTOCOL])
.timeout(ntex::time::Seconds(60))
.build(SharedCfg::default())
.await
.map_err(|e| WsConnectError::BuilderError(e.to_string()))?;
Ok(ws_client.connect().await?.seal())
} else {
Err(WsConnectError::WrongUriSchema(uri.to_string()))
}
}
type WsResponse = Result<SubgraphResponse<'static>, WsClientError>;
pub(crate) type WsResponseStream = LocalBoxStream<'static, WsResponse>;
#[derive(Clone)]
struct ClientSubscription {
sender: mpsc::Sender<WsResponse>,
custom_scalar_paths: Option<CustomScalarPaths>,
}
type WsStateRef = Rc<RefCell<WsState<ClientSubscription>>>;
pub struct Connected {
connection: WsConnection<Sealed>,
}
pub struct Initialized {
sink: ws::WsSink,
state: WsStateRef,
next_subscription_id: u64,
_heartbeat_stop_tx: Option<oneshot::Sender<()>>,
dispatcher_done_rx: Option<oneshot::Receiver<WsClientError>>,
}
pub struct WsClient<State> {
state: State,
}
impl WsClient<Connected> {
pub fn new(connection: WsConnection<Sealed>) -> Self {
Self {
state: Connected { connection },
}
}
pub async fn init(
self,
payload: Option<ConnectionInitPayload>,
) -> Result<WsClient<Initialized>, WsInitError> {
debug!(target: targets::WEBSOCKET_CLIENT, "Initialising WebSocket client connection");
let sink = self.state.connection.sink();
let mut receiver = self.state.connection.receiver();
let (acknowledged_tx, acknowledged_rx) = oneshot::channel();
let state: WsStateRef = Rc::new(RefCell::new(WsState::new(acknowledged_tx)));
let (heartbeat_stop_tx, heartbeat_stop_rx) = oneshot::channel();
rt::spawn(heartbeat(state.clone(), sink.clone(), heartbeat_stop_rx));
rt::spawn(handshake_timeout(
state.clone(),
sink.clone(),
acknowledged_rx,
CloseCode::ConnectionAcknowledgementTimeout,
));
sink.send(ClientMessage::init(payload)).await?;
loop {
match receiver.next().await {
Some(Ok(frame)) => {
match parse_frame_to_text(frame, &state) {
Ok(text) => {
let server_msg = match text_to_server_message(&text) {
Ok(msg) => msg,
Err(msg) => {
let _ = sink.send(msg).await;
return Err(WsInitError::InvalidMessage);
}
};
match server_msg {
ServerMessage::ConnectionAck {} => {
state.borrow_mut().handshake_received = true;
state.borrow_mut().complete_handshake();
debug!(target: targets::WEBSOCKET_CLIENT, "Connection acknowledged");
break;
}
ServerMessage::Ping {} => {
let _ = sink.send(ClientMessage::pong()).await;
}
ServerMessage::Pong {} => {}
_ => {
error!(target: targets::WEBSOCKET_CLIENT,
error = ?server_msg,
"Wrong message received before ConnectionAck",
);
let _ = sink.send(CloseCode::Unauthorized.into()).await;
return Err(WsInitError::WrongMessageBeforeAck);
}
}
}
Err(FrameNotParsedToText::Message(msg)) => {
let _ = sink.send(msg).await;
}
Err(FrameNotParsedToText::Closed) => {
debug!(target: targets::WEBSOCKET_CLIENT, "Connection closed before acknowledgement");
return Err(WsInitError::ConnectionClosedBeforeAck);
}
Err(FrameNotParsedToText::None) => {}
}
}
Some(Err(e)) => {
error!(target: targets::WEBSOCKET_CLIENT, error = ?e, "WebSocket receiver error during init");
return Err(WsInitError::ConnectionAckReceiverError);
}
None => {
debug!(target: targets::WEBSOCKET_CLIENT, "WebSocket receiver closed during init");
return Err(WsInitError::ConnectionAckReceiverClosed);
}
}
}
let dispatcher_state = state.clone();
let dispatcher_sink = sink.clone();
let (dispatcher_done_tx, dispatcher_done_rx) = oneshot::channel();
rt::spawn(async move {
let _guard = DispatcherGuard {
state: dispatcher_state.clone(),
};
let error = dispatch_loop(receiver, dispatcher_sink, dispatcher_state.clone()).await;
for (_, subscription) in dispatcher_state.borrow_mut().subscriptions.drain() {
let _ = subscription.sender.send(Err(error.clone()));
subscription.sender.close();
}
let _ = dispatcher_done_tx.send(error);
});
Ok(WsClient {
state: Initialized {
sink,
state,
next_subscription_id: 1,
_heartbeat_stop_tx: Some(heartbeat_stop_tx),
dispatcher_done_rx: Some(dispatcher_done_rx),
},
})
}
}
impl WsClient<Initialized> {
fn next_subscription_id(&mut self) -> Result<String, WsClientError> {
let id = self.state.next_subscription_id;
self.state.next_subscription_id = self
.state
.next_subscription_id
.checked_add(1)
.ok_or(WsClientError::SubscriptionIdExhausted)?;
Ok(id.to_string())
}
pub fn take_dispatcher_done(&mut self) -> oneshot::Receiver<WsClientError> {
self.state
.dispatcher_done_rx
.take()
.expect("dispatcher completion receiver can only be taken once")
}
pub async fn subscribe(
&mut self,
subscribe_payload: SubscribePayload,
custom_scalar_paths: Option<CustomScalarPaths>,
) -> Result<WsResponseStream, WsClientError> {
let subscribe_id = self.next_subscription_id()?;
let (tx, rx) = mpsc::channel();
self.state.state.borrow_mut().subscriptions.insert(
subscribe_id.clone(),
ClientSubscription {
sender: tx,
custom_scalar_paths,
},
);
let mut guard = SubscriptionGuard {
state: self.state.state.clone(),
sink: self.state.sink.clone(),
id: Some(subscribe_id.clone()),
send_complete: false,
};
self.state
.sink
.send(ClientMessage::subscribe(
subscribe_id.clone(),
subscribe_payload,
))
.await?;
guard.send_complete = true;
trace!(target: targets::WEBSOCKET_CLIENT, subscription_id = %subscribe_id, "Subscribe message sent");
Ok(Box::pin(async_stream::stream! {
let mut rx = rx;
let _guard = guard;
while let Some(response) = rx.next().await {
yield response;
}
}))
}
}
impl Drop for Initialized {
fn drop(&mut self) {
let sink = self.sink.clone();
rt::spawn(async move {
let _ = sink
.send(ws::Message::Close(Some(ws::CloseCode::Normal.into())))
.await;
});
}
}
struct SubscriptionGuard {
state: WsStateRef,
sink: WsSink,
id: Option<String>,
send_complete: bool,
}
impl Drop for SubscriptionGuard {
fn drop(&mut self) {
let Some(id) = self.id.take() else {
return;
};
if self.state.borrow_mut().subscriptions.remove(&id).is_some() && self.send_complete {
let sink = self.sink.clone();
rt::spawn(async move {
let _ = sink.send(ClientMessage::complete(id)).await;
});
}
}
}
async fn dispatch_loop(
mut receiver: mpsc::Receiver<Result<ws::Frame, WsError<()>>>,
sink: WsSink,
state: WsStateRef,
) -> WsClientError {
loop {
match receiver.next().await {
Some(Ok(frame)) => {
match parse_frame_to_text(frame, &state) {
Ok(text) => {
if let Some(msg) = handle_text_frame(text, &state) {
if send_and_is_closed(sink.clone(), msg).await {
return WsClientError::ConnectionClosed;
}
}
}
Err(FrameNotParsedToText::Message(msg)) => {
if send_and_is_closed(sink.clone(), msg).await {
return WsClientError::ConnectionClosed;
}
}
Err(FrameNotParsedToText::Closed) => {
return WsClientError::ConnectionClosed;
}
Err(FrameNotParsedToText::None) => {}
}
}
Some(Err(e)) => {
error!(target: targets::WEBSOCKET_CLIENT, error = ?e, "Dispatch loop WebSocket receiver error");
return WsClientError::MessageDispatcherClosed;
}
None => {
return WsClientError::MessageDispatcherClosed;
}
}
}
}
struct DispatcherGuard {
state: WsStateRef,
}
impl Drop for DispatcherGuard {
fn drop(&mut self) {
for (_, subscription) in self.state.borrow_mut().subscriptions.drain() {
let _ = subscription
.sender
.send(Err(WsClientError::ConnectionClosed));
subscription.sender.close();
}
}
}
async fn send_and_is_closed(sink: WsSink, msg: ws::Message) -> bool {
let is_close = matches!(msg, ws::Message::Close(_));
let _ = sink.send(msg).await;
is_close
}
fn handle_text_frame(text: String, state: &WsStateRef) -> Option<ws::Message> {
let server_msg = match text_to_server_message(&text) {
Ok(msg) => msg,
Err(msg) => return Some(msg),
};
trace!(target: targets::WEBSOCKET_CLIENT, type = server_msg.as_ref(), "Received server message");
match server_msg {
ServerMessage::ConnectionAck {} => {
None
}
ServerMessage::Next { id, payload } => {
if let Some(subscription) = state.borrow().subscriptions.get(&id) {
let payload_bytes = Bytes::from(sonic_rs::to_vec(&payload).unwrap_or_default());
let response = match SubgraphResponse::deserialize_from_bytes(
payload_bytes,
subscription.custom_scalar_paths.as_ref(),
) {
Ok(response) => Ok(response),
Err(e) => {
tracing::error!(target: targets::WEBSOCKET_CLIENT, error = ?e, "Failed to deserialize payload");
Err(WsClientError::FailedToDeserializePayload)
}
};
let _ = subscription.sender.send(response);
}
None
}
ServerMessage::Error { id, payload } => {
if let Some(subscription) = state.borrow_mut().subscriptions.remove(&id) {
let _ = subscription.sender.send(Ok(SubgraphResponse {
errors: Some(payload),
..Default::default()
}));
subscription.sender.close();
}
None
}
ServerMessage::Complete { id } => {
if let Some(subscription) = state.borrow_mut().subscriptions.remove(&id) {
subscription.sender.close();
}
None
}
ServerMessage::Ping {} => Some(ClientMessage::pong()),
ServerMessage::Pong {} => None,
}
}
fn text_to_server_message(text: &str) -> Result<ServerMessage, ws::Message> {
let server_msg: ServerMessage = match sonic_rs::from_str(text) {
Ok(msg) => msg,
Err(e) => {
error!(target: targets::WEBSOCKET_CLIENT, error = ?e, "Failed to parse server message to JSON");
return Err(CloseCode::BadResponse("Invalid message received from server").into());
}
};
Ok(server_msg)
}
#[cfg(test)]
mod tests {
use super::*;
use futures::StreamExt;
#[tokio::test]
async fn handle_text_frame_uses_subscription_custom_scalar_paths() {
let (ack_tx, _ack_rx) = oneshot::channel();
let state: WsStateRef = Rc::new(RefCell::new(WsState::new(ack_tx)));
let (tx, mut rx) = mpsc::channel();
let mut custom_scalar_paths = CustomScalarPaths::default();
custom_scalar_paths.insert_path(["custom"]);
state.borrow_mut().subscriptions.insert(
"1".to_string(),
ClientSubscription {
sender: tx,
custom_scalar_paths: Some(custom_scalar_paths),
},
);
let text =
r#"{"type":"next","id":"1","payload":{"data":{"custom":{"escaped.key\t":"value"}}}}"#
.to_string();
assert!(handle_text_frame(text, &state).is_none());
let response = rx.next().await.expect("response").expect("valid response");
let data = response.data.as_object().unwrap();
assert!(data[0].1.as_raw_json().is_some());
}
}