use std::{
sync::{Arc, Weak},
time::Duration,
};
use crate::query_planner::planner::plan_nodes::CustomScalarPaths;
use crate::telemetry::{
logging::targets,
metrics::{
subscription_metrics::SubscriptionTransport,
websocket_pool_metrics::{WebSocketPoolConnectionCloseReason, WebSocketPoolOperationType},
},
TelemetryContext,
};
use async_trait::async_trait;
use dashmap::{mapref::entry::Entry, DashMap};
use futures::{stream::BoxStream, StreamExt};
use http::{HeaderMap, Uri};
use ntex::rt;
use tokio::{
sync::{mpsc, oneshot},
time::Instant,
};
use tracing::{debug, info, trace, warn};
use crate::executor::{
executors::{
common::{ConnectionFingerprint, SubgraphExecutionRequest, SubgraphExecutor},
error::SubgraphExecutorError,
graphql_transport_ws::SubscribePayload,
subscription_buffer::drain_into,
websocket_client::{self, WsClient, WsClientError},
},
plugin_context::PluginRequestState,
response::subgraph_response::SubgraphResponse,
};
type SubscriptionItem = Result<SubgraphResponse<'static>, SubgraphExecutorError>;
type InitResult = Result<Arc<PooledWebSocketExecutor>, PoolInitError>;
type PoolEntries = DashMap<WebSocketConnectionId, PoolEntry>;
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub struct WebSocketConnectionId {
subgraph_name: Arc<str>,
endpoint: Uri,
fingerprint: ConnectionFingerprint,
}
impl WebSocketConnectionId {
pub fn new(
subgraph_name: impl Into<Arc<str>>,
endpoint: Uri,
fingerprint: ConnectionFingerprint,
) -> Self {
Self {
subgraph_name: subgraph_name.into(),
endpoint,
fingerprint,
}
}
}
#[derive(Clone, Debug, thiserror::Error)]
enum PoolInitError {
#[error("WebSocket connection failed: {0}")]
Connect(String),
#[error("WebSocket handshake failed: {0}")]
Handshake(String),
}
impl PoolInitError {
fn into_executor_error(self, endpoint: &Uri) -> SubgraphExecutorError {
match self {
Self::Connect(error) => {
SubgraphExecutorError::WebSocketConnectFailure(endpoint.to_string(), error)
}
Self::Handshake(error) => {
SubgraphExecutorError::WebSocketHandshakeFailure(endpoint.to_string(), error)
}
}
}
}
struct ConnectingEntry {
generation: Arc<()>,
waiters: Vec<oneshot::Sender<InitResult>>,
}
enum PoolEntry {
Connecting(ConnectingEntry),
Initialized(Arc<PooledWebSocketExecutor>),
}
pub struct WebSocketInit {
pub endpoint: Uri,
pub headers: HeaderMap,
pub tls_config: Option<Arc<rustls::ClientConfig>>,
pub buffer_capacity: usize,
pub idle_timeout: Duration,
pub telemetry_context: Arc<TelemetryContext>,
}
#[derive(Default)]
pub struct WebSocketPool {
entries: Arc<PoolEntries>,
}
impl WebSocketPool {
pub fn get_initialized(
&self,
id: &WebSocketConnectionId,
) -> Option<Arc<PooledWebSocketExecutor>> {
let executor = self.entries.get(id).and_then(|entry| match entry.value() {
PoolEntry::Initialized(executor) if !executor.commands.is_closed() => {
Some(executor.clone())
}
PoolEntry::Connecting(_) | PoolEntry::Initialized(_) => None,
});
trace!(
target: targets::WEBSOCKET_POOL,
subgraph = %id.subgraph_name,
found = executor.is_some(),
"looked up WebSocket pool connection"
);
executor
}
pub async fn get_or_initialize(
&self,
id: WebSocketConnectionId,
init: WebSocketInit,
) -> Result<Arc<PooledWebSocketExecutor>, SubgraphExecutorError> {
let endpoint = init.endpoint.clone();
let (wait_rx, generation) = match self.entries.entry(id.clone()) {
Entry::Occupied(mut entry) => match entry.get_mut() {
PoolEntry::Initialized(executor) if !executor.commands.is_closed() => {
trace!(
target: targets::WEBSOCKET_POOL,
subgraph = %id.subgraph_name,
endpoint = %endpoint,
"reusing WebSocket pool connection"
);
return Ok(executor.clone());
}
PoolEntry::Connecting(connecting) => {
connecting.waiters.retain(|waiter| !waiter.is_closed());
let (wait_tx, wait_rx) = oneshot::channel();
connecting.waiters.push(wait_tx);
(wait_rx, None)
}
PoolEntry::Initialized(_) => {
let generation = Arc::new(());
let (wait_tx, wait_rx) = oneshot::channel();
entry.insert(PoolEntry::Connecting(ConnectingEntry {
generation: generation.clone(),
waiters: vec![wait_tx],
}));
(wait_rx, Some(generation))
}
},
Entry::Vacant(entry) => {
let generation = Arc::new(());
let (wait_tx, wait_rx) = oneshot::channel();
entry.insert(PoolEntry::Connecting(ConnectingEntry {
generation: generation.clone(),
waiters: vec![wait_tx],
}));
(wait_rx, Some(generation))
}
};
if let Some(generation) = generation {
debug!(
target: targets::WEBSOCKET_POOL,
subgraph = %id.subgraph_name,
endpoint = %endpoint,
"initializing WebSocket pool connection"
);
let cleanup = InitializationCleanup::new(self.entries.clone(), id.clone(), generation);
let telemetry_context = init.telemetry_context.clone();
let log_endpoint = endpoint.clone();
rt::spawn(async move {
let mut cleanup = cleanup;
match initialize_connection(&cleanup.entries, &cleanup.id, init).await {
Ok(connection) => {
telemetry_context
.metrics
.websocket_pool
.record_connection_initialization(&cleanup.id.subgraph_name, true);
let executor = connection.executor.clone();
let Some(waiters) = cleanup.publish(executor.clone()) else {
debug!(
target: targets::WEBSOCKET_POOL,
subgraph = %cleanup.id.subgraph_name,
endpoint = %log_endpoint,
"discarding stale WebSocket pool connection initialization"
);
return;
};
info!(
target: targets::WEBSOCKET_POOL,
subgraph = %cleanup.id.subgraph_name,
endpoint = %log_endpoint,
waiting_requests = waiters.len(),
"WebSocket pool connection initialized"
);
rt::spawn(connection.owner.run());
notify_waiters(waiters, Ok(executor));
}
Err(error) => {
telemetry_context
.metrics
.websocket_pool
.record_connection_initialization(&cleanup.id.subgraph_name, false);
warn!(
target: targets::WEBSOCKET_POOL,
subgraph = %cleanup.id.subgraph_name,
endpoint = %log_endpoint,
error = %error,
"failed to initialize WebSocket pool connection"
);
let waiters = cleanup.remove();
notify_waiters(waiters, Err(error));
}
}
});
} else {
debug!(
target: targets::WEBSOCKET_POOL,
subgraph = %id.subgraph_name,
endpoint = %endpoint,
"waiting for WebSocket pool connection initialization"
);
init.telemetry_context
.metrics
.websocket_pool
.record_connection_initialization_waiter(&id.subgraph_name);
}
wait_rx
.await
.map_err(|_| {
warn!(
target: targets::WEBSOCKET_POOL,
subgraph = %id.subgraph_name,
endpoint = %endpoint,
"WebSocket pool connection initialization stopped before completion"
);
SubgraphExecutorError::WebSocketArbiterChannelClosed
})?
.map_err(|error| error.into_executor_error(&endpoint))
}
}
fn notify_waiters(waiters: Vec<oneshot::Sender<InitResult>>, result: InitResult) {
for waiter in waiters {
let _ = waiter.send(result.clone());
}
}
struct InitializationCleanup {
entries: Arc<PoolEntries>,
id: WebSocketConnectionId,
generation: Arc<()>,
armed: bool,
}
impl InitializationCleanup {
fn new(entries: Arc<PoolEntries>, id: WebSocketConnectionId, generation: Arc<()>) -> Self {
Self {
entries,
id,
generation,
armed: true,
}
}
fn publish(
&mut self,
executor: Arc<PooledWebSocketExecutor>,
) -> Option<Vec<oneshot::Sender<InitResult>>> {
let waiters = match self.entries.entry(self.id.clone()) {
Entry::Occupied(mut entry) => {
let PoolEntry::Connecting(connecting) = entry.get_mut() else {
return None;
};
if !Arc::ptr_eq(&connecting.generation, &self.generation) {
return None;
}
let waiters = std::mem::take(&mut connecting.waiters);
entry.insert(PoolEntry::Initialized(executor));
waiters
}
Entry::Vacant(_) => return None,
};
self.armed = false;
Some(waiters)
}
fn remove(&mut self) -> Vec<oneshot::Sender<InitResult>> {
let removed = self.entries.remove_if(&self.id, |_, entry| {
matches!(
entry,
PoolEntry::Connecting(connecting)
if Arc::ptr_eq(&connecting.generation, &self.generation)
)
});
self.armed = false;
match removed {
Some((_, PoolEntry::Connecting(connecting))) => connecting.waiters,
Some((_, PoolEntry::Initialized(_))) | None => Vec::new(),
}
}
}
impl Drop for InitializationCleanup {
fn drop(&mut self) {
if self.armed {
let _ = self.entries.remove_if(&self.id, |_, entry| {
matches!(
entry,
PoolEntry::Connecting(connecting)
if Arc::ptr_eq(&connecting.generation, &self.generation)
)
});
}
}
}
struct InitializedConnection {
executor: Arc<PooledWebSocketExecutor>,
owner: ConnectionOwner,
}
async fn initialize_connection(
entries: &Arc<PoolEntries>,
id: &WebSocketConnectionId,
init: WebSocketInit,
) -> Result<InitializedConnection, PoolInitError> {
let wsconn = websocket_client::connect(&init.endpoint, init.tls_config)
.await
.map_err(|error| PoolInitError::Connect(error.to_string()))?;
let client = WsClient::new(wsconn);
let init_payload = (!init.headers.is_empty()).then(|| init.headers.into());
let mut client = client
.init(init_payload)
.await
.map_err(|error| PoolInitError::Handshake(error.to_string()))?;
let dispatcher_done = client.take_dispatcher_done();
let (commands, task_commands) = mpsc::channel(init.buffer_capacity);
let endpoint = Arc::<str>::from(init.endpoint.to_string());
let executor = Arc::new(PooledWebSocketExecutor {
commands,
telemetry_context: init.telemetry_context.clone(),
buffer_capacity: init.buffer_capacity,
entries: Arc::downgrade(entries),
id: id.clone(),
endpoint_uri: init.endpoint,
endpoint: endpoint.clone(),
});
let owner = ConnectionOwner {
client,
dispatcher_done,
commands: task_commands,
executor: Arc::downgrade(&executor),
telemetry_context: init.telemetry_context,
subgraph_name: id.subgraph_name.clone(),
endpoint,
idle_timeout: init.idle_timeout,
};
Ok(InitializedConnection { executor, owner })
}
struct ConnectionCommand {
payload: SubscribePayload,
custom_scalar_paths: Option<CustomScalarPaths>,
responses: mpsc::Sender<SubscriptionItem>,
ready: oneshot::Sender<Result<(), WsClientError>>,
}
pub struct PooledWebSocketExecutor {
commands: mpsc::Sender<ConnectionCommand>,
telemetry_context: Arc<TelemetryContext>,
buffer_capacity: usize,
entries: Weak<PoolEntries>,
id: WebSocketConnectionId,
endpoint_uri: Uri,
endpoint: Arc<str>,
}
impl PooledWebSocketExecutor {
fn evict_if_current(&self) {
if let Some(entries) = self.entries.upgrade() {
if entries
.remove_if(&self.id, |_, entry| {
matches!(
entry,
PoolEntry::Initialized(current)
if current.commands.same_channel(&self.commands)
)
})
.is_some()
{
debug!(
target: targets::WEBSOCKET_POOL,
subgraph = %self.id.subgraph_name,
endpoint = %self.endpoint,
"evicted WebSocket pool connection"
);
}
}
}
async fn submit(
&self,
execution_request: SubgraphExecutionRequest<'_>,
response_capacity: usize,
) -> Result<mpsc::Receiver<SubscriptionItem>, SubgraphExecutorError> {
let permit = self.commands.reserve().await.map_err(|_| {
debug!(
target: targets::WEBSOCKET_POOL,
subgraph = %self.id.subgraph_name,
endpoint = %self.endpoint,
"WebSocket pool connection closed before operation could be queued"
);
self.evict_if_current();
SubgraphExecutorError::WebSocketArbiterChannelClosed
})?;
let custom_scalar_paths = execution_request.custom_scalar_paths.cloned();
let payload = SubscribePayload::try_from(execution_request)?;
let (responses, receiver) = mpsc::channel(response_capacity);
let (ready, ready_rx) = oneshot::channel();
permit.send(ConnectionCommand {
payload,
custom_scalar_paths,
responses,
ready,
});
match ready_rx.await {
Ok(Ok(())) => Ok(receiver),
Ok(Err(error)) => {
warn!(
target: targets::WEBSOCKET_POOL,
subgraph = %self.id.subgraph_name,
endpoint = %self.endpoint,
error = %error,
"failed to start operation on WebSocket pool connection"
);
Err(error.into())
}
Err(_) => {
debug!(
target: targets::WEBSOCKET_POOL,
subgraph = %self.id.subgraph_name,
endpoint = %self.endpoint,
"WebSocket pool connection closed while starting operation"
);
self.evict_if_current();
Err(SubgraphExecutorError::WebSocketArbiterChannelClosed)
}
}
}
}
#[async_trait]
impl SubgraphExecutor for PooledWebSocketExecutor {
fn executor_name(&self) -> &str {
"pooled-websocket"
}
fn endpoint(&self) -> &Uri {
&self.endpoint_uri
}
async fn execute<'a>(
&self,
execution_request: SubgraphExecutionRequest<'a>,
timeout: Option<Duration>,
_plugin_req_state: Option<&'a PluginRequestState<'a>>,
) -> Result<SubgraphResponse<'static>, SubgraphExecutorError> {
let _operation_guard = self
.telemetry_context
.metrics
.websocket_pool
.active_operation(&self.id.subgraph_name, WebSocketPoolOperationType::Execute);
let operation = async {
let mut responses = self.submit(execution_request, 1).await?;
responses.recv().await.ok_or_else(|| {
SubgraphExecutorError::WebSocketStreamClosedEmpty(self.endpoint.to_string())
})?
};
match timeout {
Some(timeout) => tokio::time::timeout(timeout, operation).await?,
None => operation.await,
}
}
async fn subscribe<'a>(
&self,
execution_request: SubgraphExecutionRequest<'a>,
_timeout: Option<Duration>,
) -> Result<BoxStream<'static, SubscriptionItem>, SubgraphExecutorError> {
let pool_operation_guard = self
.telemetry_context
.metrics
.websocket_pool
.active_operation(
&self.id.subgraph_name,
WebSocketPoolOperationType::Subscribe,
);
let mut responses = self.submit(execution_request, self.buffer_capacity).await?;
let subscription_operation_guard = self
.telemetry_context
.metrics
.subscriptions
.active_subgraph_operation(&self.id.subgraph_name);
Ok(Box::pin(async_stream::stream! {
let _pool_operation_guard = pool_operation_guard;
let _subscription_operation_guard = subscription_operation_guard;
while let Some(item) = responses.recv().await {
yield item;
}
}))
}
}
#[async_trait]
impl SubgraphExecutor for Arc<PooledWebSocketExecutor> {
fn executor_name(&self) -> &str {
self.as_ref().executor_name()
}
fn endpoint(&self) -> &Uri {
self.as_ref().endpoint()
}
async fn execute<'a>(
&self,
execution_request: SubgraphExecutionRequest<'a>,
timeout: Option<Duration>,
plugin_req_state: Option<&'a PluginRequestState<'a>>,
) -> Result<SubgraphResponse<'static>, SubgraphExecutorError> {
self.as_ref()
.execute(execution_request, timeout, plugin_req_state)
.await
}
async fn subscribe<'a>(
&self,
execution_request: SubgraphExecutionRequest<'a>,
timeout: Option<Duration>,
) -> Result<BoxStream<'static, SubscriptionItem>, SubgraphExecutorError> {
self.as_ref().subscribe(execution_request, timeout).await
}
}
struct OperationCompletionGuard(mpsc::UnboundedSender<()>);
impl Drop for OperationCompletionGuard {
fn drop(&mut self) {
let _ = self.0.send(());
}
}
enum ConnectionShutdown {
Idle,
Dispatcher(WsClientError),
CommandsClosed,
}
impl ConnectionShutdown {
fn command_error(&self) -> WsClientError {
match self {
Self::Idle => WsClientError::ConnectionClosed,
Self::Dispatcher(error) => error.clone(),
Self::CommandsClosed => WsClientError::MessageDispatcherClosed,
}
}
}
struct ConnectionOwner {
client: WsClient<crate::executor::executors::websocket_client::Initialized>,
dispatcher_done: ntex::channel::oneshot::Receiver<WsClientError>,
commands: mpsc::Receiver<ConnectionCommand>,
executor: Weak<PooledWebSocketExecutor>,
telemetry_context: Arc<TelemetryContext>,
subgraph_name: Arc<str>,
endpoint: Arc<str>,
idle_timeout: Duration,
}
impl ConnectionOwner {
fn evict(&self) {
if let Some(executor) = self.executor.upgrade() {
executor.evict_if_current();
}
}
async fn run(mut self) {
let telemetry_context = self.telemetry_context.clone();
let subgraph_name = self.subgraph_name.clone();
let _connection_guard = telemetry_context
.metrics
.websocket_pool
.active_connection(&subgraph_name);
let (completed_tx, mut completed_rx) = mpsc::unbounded_channel();
let mut active_operations = 0usize;
let idle_timer = tokio::time::sleep(self.idle_timeout);
tokio::pin!(idle_timer);
let shutdown = 'connection: loop {
tokio::select! {
command = self.commands.recv() => {
let Some(ConnectionCommand {
payload,
custom_scalar_paths,
responses,
ready,
}) = command else {
break ConnectionShutdown::CommandsClosed;
};
if responses.is_closed() {
continue;
}
if active_operations == 0 {
idle_timer
.as_mut()
.reset(Instant::now() + self.idle_timeout);
}
let subscribe_result = {
let subscribe = self.client.subscribe(payload, custom_scalar_paths);
tokio::pin!(subscribe);
tokio::select! {
result = &mut subscribe => Some(result),
_ = responses.closed() => None,
dispatcher = &mut self.dispatcher_done => {
let error = dispatcher
.unwrap_or(WsClientError::MessageDispatcherClosed);
let _ = ready.send(Err(error.clone()));
break 'connection ConnectionShutdown::Dispatcher(error);
}
}
};
let Some(subscribe_result) = subscribe_result else {
continue;
};
match subscribe_result {
Ok(stream) => {
if ready.send(Ok(())).is_err() {
drop(stream);
continue;
}
active_operations += 1;
let completion_guard =
OperationCompletionGuard(completed_tx.clone());
let telemetry_context = self.telemetry_context.clone();
let subgraph_name = self.subgraph_name.clone();
let endpoint = self.endpoint.clone();
rt::spawn(async move {
let _completion_guard = completion_guard;
drain_into(
stream.map(|item| item.map_err(SubgraphExecutorError::from)),
responses,
&telemetry_context,
SubscriptionTransport::WebSocket,
&subgraph_name,
&endpoint,
)
.await;
});
}
Err(error) => {
let _ = ready.send(Err(error));
}
}
}
Some(()) = completed_rx.recv(), if active_operations > 0 => {
active_operations -= 1;
if active_operations == 0 {
idle_timer
.as_mut()
.reset(Instant::now() + self.idle_timeout);
}
}
dispatcher = &mut self.dispatcher_done => {
break ConnectionShutdown::Dispatcher(
dispatcher.unwrap_or(WsClientError::MessageDispatcherClosed),
);
}
_ = &mut idle_timer, if active_operations == 0 => {
break ConnectionShutdown::Idle;
}
}
};
self.commands.close();
self.evict();
while let Some(command) = self.commands.recv().await {
let _ = command.ready.send(Err(shutdown.command_error()));
}
let reason = match shutdown {
ConnectionShutdown::Idle => {
info!(
target: targets::WEBSOCKET_POOL,
subgraph = %self.subgraph_name,
endpoint = %self.endpoint,
"closing idle WebSocket pool connection"
);
WebSocketPoolConnectionCloseReason::Idle
}
ConnectionShutdown::Dispatcher(error) => {
warn!(
target: targets::WEBSOCKET_POOL,
subgraph = %self.subgraph_name,
endpoint = %self.endpoint,
error = %error,
"WebSocket pool connection dispatcher stopped"
);
WebSocketPoolConnectionCloseReason::Dispatcher
}
ConnectionShutdown::CommandsClosed => {
info!(
target: targets::WEBSOCKET_POOL,
subgraph = %self.subgraph_name,
endpoint = %self.endpoint,
"WebSocket pool dropped, closing connection"
);
WebSocketPoolConnectionCloseReason::PoolDropped
}
};
self.telemetry_context
.metrics
.websocket_pool
.record_connection_closed(&self.subgraph_name, reason);
}
}
impl Drop for ConnectionOwner {
fn drop(&mut self) {
self.commands.close();
self.evict();
}
}