use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
use rand::{Rng, thread_rng};
use tokio::sync::{Mutex, RwLock, Semaphore, mpsc, oneshot};
use crate::connector::{BaseConnector, ConnectorConfig};
use crate::error::{ConnectorError, Result};
use crate::logger::Logger;
use crate::multi::{MultiTransportOptions, RegistrationKey};
use crate::transport::{Transport, TransportOptions, TransportType, WebSocketTransport};
use crate::types::{ConnectorMetrics, ExecuteRequest as SdkExecuteRequest, PayloadEncoding};
use crate::utils::{deserialize_payload, error_response, sanitize_identifier, serialize_payload};
use strike48_proto::proto::{
self, ConnectorCapabilities, HeartbeatRequest, HeartbeatResponse, InstanceMetadata,
RegisterConnectorRequest, StreamMessage, stream_message,
};
const SHARED_OUTBOUND_CAPACITY: usize = 256;
const SHUTDOWN_POLL: Duration = Duration::from_millis(50);
const DEFAULT_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(30);
struct HandleEntry {
key: RegistrationKey,
inbound_tx: mpsc::Sender<StreamMessage>,
pending_register: Mutex<Option<oneshot::Sender<proto::RegisterConnectorResponse>>>,
}
pub(crate) struct WsMultiplexSocket {
pub tenant_id: String,
outbound_tx: mpsc::Sender<StreamMessage>,
by_arn: RwLock<HashMap<String, Arc<HandleEntry>>>,
by_dot: RwLock<HashMap<String, Arc<HandleEntry>>>,
by_request_id: RwLock<HashMap<String, Arc<HandleEntry>>>,
shutdown: Arc<AtomicBool>,
last_inbound: Mutex<Instant>,
max_handlers: usize,
register_acked: Arc<AtomicBool>,
}
impl WsMultiplexSocket {
fn new(
tenant_id: String,
outbound_tx: mpsc::Sender<StreamMessage>,
max_handlers: usize,
) -> Self {
Self {
tenant_id,
outbound_tx,
by_arn: RwLock::new(HashMap::new()),
by_dot: RwLock::new(HashMap::new()),
by_request_id: RwLock::new(HashMap::new()),
shutdown: Arc::new(AtomicBool::new(false)),
last_inbound: Mutex::new(Instant::now()),
max_handlers: max_handlers.max(1),
register_acked: Arc::new(AtomicBool::new(false)),
}
}
async fn admit(
self: &Arc<Self>,
key: RegistrationKey,
inbound_tx: mpsc::Sender<StreamMessage>,
pending_register: oneshot::Sender<proto::RegisterConnectorResponse>,
) -> Result<Arc<HandleEntry>> {
let dot = key.to_string();
let mut by_dot = self.by_dot.write().await;
if by_dot.len() >= self.max_handlers {
return Err(ConnectorError::InvalidConfig(format!(
"ws multiplex: cannot admit registration {dot}; \
per-socket handler cap of {} reached for tenant '{}'",
self.max_handlers, self.tenant_id
)));
}
if by_dot.contains_key(&dot) {
return Err(ConnectorError::InvalidConfig(format!(
"ws multiplex: registration {dot} is already admitted on this socket"
)));
}
let entry = Arc::new(HandleEntry {
key,
inbound_tx,
pending_register: Mutex::new(Some(pending_register)),
});
by_dot.insert(dot, entry.clone());
Ok(entry)
}
async fn evict(&self, key: &RegistrationKey, arn: Option<&str>) {
let dot = key.to_string();
self.by_dot.write().await.remove(&dot);
if let Some(arn) = arn {
self.by_arn.write().await.remove(arn);
}
let mut by_id = self.by_request_id.write().await;
by_id.retain(|_, h| h.key != *key);
}
async fn sole_handle(&self) -> Option<Arc<HandleEntry>> {
let by_dot = self.by_dot.read().await;
if by_dot.len() == 1 {
by_dot.values().next().cloned()
} else {
None
}
}
async fn unique_by_instance_id(&self, target: &str) -> Option<Arc<HandleEntry>> {
if target.is_empty() {
return None;
}
let by_dot = self.by_dot.read().await;
let mut matches = by_dot
.values()
.filter(|entry| entry.key.instance_id == target);
let first = matches.next()?.clone();
if matches.next().is_some() {
return None;
}
Some(first)
}
async fn bind_arn(&self, arn: String, handle: Arc<HandleEntry>) {
self.by_arn.write().await.insert(arn, handle);
}
#[allow(dead_code)]
async fn track_request_id(&self, request_id: String, handle: Arc<HandleEntry>) {
self.by_request_id.write().await.insert(request_id, handle);
}
#[allow(dead_code)]
fn outbound(&self) -> mpsc::Sender<StreamMessage> {
self.outbound_tx.clone()
}
}
async fn open_socket(
tenant_id: &str,
opts: &MultiTransportOptions,
runner_shutdown: Arc<AtomicBool>,
) -> Result<(Arc<WsMultiplexSocket>, tokio::task::JoinHandle<()>)> {
let transport_opts = TransportOptions {
host: opts.host.clone(),
use_tls: opts.use_tls,
connect_timeout_ms: Some(opts.connect_timeout_ms),
default_timeout_ms: None,
channel_capacity: Some(SHARED_OUTBOUND_CAPACITY),
};
let mut transport = WebSocketTransport::new(transport_opts);
let (writer_tx, mut reader_rx) = transport.start_stream(None).await?;
let (outbound_tx, mut outbound_rx) = mpsc::channel::<StreamMessage>(SHARED_OUTBOUND_CAPACITY);
let socket = Arc::new(WsMultiplexSocket::new(
tenant_id.to_string(),
outbound_tx,
opts.max_handlers_per_socket,
));
let bridge_handle = {
let writer_tx = writer_tx.clone();
tokio::spawn(async move {
while let Some(msg) = outbound_rx.recv().await {
if writer_tx.send(msg).is_err() {
break;
}
}
})
};
let heartbeat_handle = {
let socket_for_hb = socket.clone();
let runner_shutdown = runner_shutdown.clone();
let interval_dur = opts
.heartbeat_interval
.unwrap_or(DEFAULT_HEARTBEAT_INTERVAL);
tokio::spawn(async move {
let mut tick = tokio::time::interval(interval_dur);
tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
tick.tick().await;
loop {
tokio::select! {
_ = tick.tick() => {}
_ = tokio::time::sleep(SHUTDOWN_POLL) => {
if runner_shutdown.load(Ordering::SeqCst)
|| socket_for_hb.shutdown.load(Ordering::SeqCst)
{
return;
}
continue;
}
}
if runner_shutdown.load(Ordering::SeqCst)
|| socket_for_hb.shutdown.load(Ordering::SeqCst)
{
return;
}
let hb = StreamMessage {
message: Some(stream_message::Message::HeartbeatRequest(
HeartbeatRequest {
gateway_id: String::new(),
timestamp_ms: now_ms(),
},
)),
};
if socket_for_hb.outbound_tx.send(hb).await.is_err() {
return;
}
}
})
};
let socket_clone = socket.clone();
let heartbeat_for_demux = heartbeat_handle.abort_handle();
let demux = tokio::spawn(async move {
let logger = Logger::new("multi/ws");
let _transport_keepalive = transport;
while let Some(msg) = reader_rx.recv().await {
*socket_clone.last_inbound.lock().await = Instant::now();
if runner_shutdown.load(Ordering::SeqCst)
|| socket_clone.shutdown.load(Ordering::SeqCst)
{
break;
}
dispatch_inbound_to_handle(&socket_clone, msg, &logger).await;
}
socket_clone.shutdown.store(true, Ordering::SeqCst);
let mut by_dot = socket_clone.by_dot.write().await;
by_dot.clear();
socket_clone.by_arn.write().await.clear();
socket_clone.by_request_id.write().await.clear();
bridge_handle.abort();
heartbeat_for_demux.abort();
});
Ok((socket, demux))
}
async fn dispatch_inbound_to_handle(
socket: &Arc<WsMultiplexSocket>,
msg: StreamMessage,
logger: &Logger,
) {
match msg.message.as_ref() {
Some(stream_message::Message::HeartbeatRequest(_)) => {
let resp = StreamMessage {
message: Some(stream_message::Message::HeartbeatResponse(
HeartbeatResponse {
gateway_id: String::new(),
timestamp_ms: now_ms(),
should_reconnect: false,
},
)),
};
let _ = socket.outbound_tx.send(resp).await;
}
Some(stream_message::Message::HeartbeatResponse(_)) => {
}
Some(stream_message::Message::RegisterResponse(resp)) => {
let parsed = arn_parts(&resp.connector_arn);
let dot = parsed.as_ref().map(|p| format!("{}.{}.{}", p.0, p.1, p.2));
let instance_from_arn = parsed.as_ref().map(|p| p.2);
let mut handle: Option<Arc<HandleEntry>> = None;
if let Some(d) = &dot {
handle = socket.by_dot.read().await.get(d).cloned();
}
if handle.is_none()
&& let Some(inst) = instance_from_arn
{
handle = socket.unique_by_instance_id(inst).await;
if handle.is_some() {
logger.debug(&format!(
"ws multiplex: register_response for arn '{}' \
routed by instance_id fallback (tenant '{}')",
resp.connector_arn, socket.tenant_id
));
}
}
if handle.is_none() {
handle = socket.sole_handle().await;
}
let handle = match handle {
Some(h) => h,
None => {
if dot.is_none() {
logger.warn(&format!(
"ws multiplex: register_response with malformed arn '{}' \
dropped (tenant '{}')",
resp.connector_arn, socket.tenant_id
));
} else {
logger.warn(&format!(
"ws multiplex: register_response for unknown arn '{}' \
dropped (tenant '{}')",
resp.connector_arn, socket.tenant_id
));
}
return;
}
};
if resp.success && !resp.connector_arn.is_empty() {
socket
.bind_arn(resp.connector_arn.clone(), handle.clone())
.await;
}
let mut pending = handle.pending_register.lock().await;
if let Some(tx) = pending.take() {
let _ = tx.send(resp.clone());
return;
}
drop(pending);
forward_to_handle(&handle, msg, logger).await;
}
Some(stream_message::Message::ExecuteRequest(req)) => {
let arn = req.context.get("connector_arn").map(String::as_str);
route_inbound_by_arn(socket, msg.clone(), arn, "execute_request", logger).await;
}
Some(stream_message::Message::InvokeRequest(req)) => {
let arn = req.context.get("connector_arn").map(String::as_str);
route_inbound_by_arn(socket, msg.clone(), arn, "invoke_request", logger).await;
}
Some(stream_message::Message::ExecuteResponse(resp)) => {
route_inbound_by_request_id(
socket,
msg.clone(),
&resp.request_id,
"execute_response",
logger,
)
.await;
}
Some(stream_message::Message::InvokeResponse(resp)) => {
route_inbound_by_request_id(
socket,
msg.clone(),
&resp.request_id,
"invoke_response",
logger,
)
.await;
}
Some(stream_message::Message::CredentialsIssued(creds)) => {
let gid = creds.gateway_id.clone();
route_inbound_by_gateway_id(socket, msg.clone(), &gid, "credentials_issued", logger)
.await;
}
Some(stream_message::Message::ApprovalNotification(notif)) => {
let gid = notif.gateway_id.clone();
route_inbound_by_gateway_id(socket, msg.clone(), &gid, "approval_notification", logger)
.await;
}
Some(other) => {
if let Some(handle) = socket.sole_handle().await {
forward_to_handle(&handle, msg.clone(), logger).await;
} else {
logger.debug(&format!(
"ws multiplex: unrouted inbound variant {:?} dropped (tenant '{}')",
std::mem::discriminant(other),
socket.tenant_id
));
}
}
None => {
logger.debug("ws multiplex: empty inbound message");
}
}
}
async fn route_inbound_by_arn(
socket: &Arc<WsMultiplexSocket>,
msg: StreamMessage,
arn: Option<&str>,
label: &'static str,
logger: &Logger,
) {
if let Some(arn) = arn {
let map = socket.by_arn.read().await;
if let Some(handle) = map.get(arn).cloned() {
drop(map);
forward_to_handle(&handle, msg, logger).await;
return;
}
logger.warn(&format!(
"ws multiplex: {label} with unmatched connector_arn={arn:?} dropped (tenant '{}', \
{} handles admitted)",
socket.tenant_id,
socket.by_dot.read().await.len()
));
return;
}
if let Some(handle) = socket.sole_handle().await {
forward_to_handle(&handle, msg, logger).await;
return;
}
logger.warn(&format!(
"ws multiplex: {label} without connector_arn dropped (tenant '{}', \
{} handles admitted)",
socket.tenant_id,
socket.by_dot.read().await.len()
));
}
async fn route_inbound_by_request_id(
socket: &Arc<WsMultiplexSocket>,
msg: StreamMessage,
request_id: &str,
label: &'static str,
logger: &Logger,
) {
if !request_id.is_empty() {
if let Some(handle) = socket.by_request_id.write().await.remove(request_id) {
forward_to_handle(&handle, msg, logger).await;
return;
}
logger.warn(&format!(
"ws multiplex: {label} request_id={request_id:?} matched no outstanding handle \
(tenant '{}')",
socket.tenant_id
));
return;
}
if let Some(handle) = socket.sole_handle().await {
forward_to_handle(&handle, msg, logger).await;
return;
}
logger.warn(&format!(
"ws multiplex: {label} without request_id dropped (tenant '{}')",
socket.tenant_id
));
}
async fn route_inbound_by_gateway_id(
socket: &Arc<WsMultiplexSocket>,
msg: StreamMessage,
gateway_id: &str,
label: &'static str,
logger: &Logger,
) {
if !gateway_id.is_empty() {
let map = socket.by_dot.read().await;
if let Some(handle) = map.get(gateway_id).cloned() {
drop(map);
forward_to_handle(&handle, msg, logger).await;
return;
}
logger.warn(&format!(
"ws multiplex: {label} for unknown gateway_id={gateway_id:?} dropped (tenant '{}', \
{} handles admitted)",
socket.tenant_id,
socket.by_dot.read().await.len()
));
return;
}
if let Some(handle) = socket.sole_handle().await {
forward_to_handle(&handle, msg, logger).await;
return;
}
logger.warn(&format!(
"ws multiplex: {label} without gateway_id dropped (tenant '{}', \
{} handles admitted; cannot disambiguate)",
socket.tenant_id,
socket.by_dot.read().await.len()
));
}
async fn forward_to_handle(handle: &Arc<HandleEntry>, msg: StreamMessage, logger: &Logger) {
if handle.inbound_tx.send(msg).await.is_err() {
logger.debug(&format!(
"ws multiplex: handle {} inbound channel closed; frame dropped",
handle.key
));
}
}
fn arn_parts(arn: &str) -> Option<(&str, &str, &str)> {
let parts: Vec<&str> = arn.splitn(4, ':').collect();
if parts.len() < 4 || parts[0] != "matrix" {
return None;
}
Some((parts[1], parts[2], parts[3]))
}
fn now_ms() -> i64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as i64
}
pub(crate) struct WsMultiplexHandle {
pub key: RegistrationKey,
pub config: Arc<RwLock<ConnectorConfig>>,
pub connector: Arc<dyn BaseConnector>,
pub socket: Arc<WsMultiplexSocket>,
pub shutdown: Arc<AtomicBool>,
pub metrics: Arc<Mutex<ConnectorMetrics>>,
pub opts: MultiTransportOptions,
pub request_semaphore: Arc<Semaphore>,
pub session_token: Arc<RwLock<Option<String>>>,
}
impl WsMultiplexHandle {
pub async fn run(self) -> Result<()> {
let logger = Logger::new("multi/ws/handle");
if self.shutdown.load(Ordering::SeqCst) {
return Ok(());
}
let mut attempt: u32 = 0;
let mut current_arn: Option<String> = None;
let mut ever_registered = false;
loop {
if self.shutdown.load(Ordering::SeqCst) {
self.socket.evict(&self.key, current_arn.as_deref()).await;
return Ok(());
}
if self.socket.shutdown.load(Ordering::SeqCst) {
logger.debug(&format!(
"ws multiplex handle {}: socket shut down, exiting",
self.key
));
self.socket.evict(&self.key, current_arn.as_deref()).await;
return Ok(());
}
let (inbound_tx, inbound_rx) = mpsc::channel::<StreamMessage>(64);
let (register_tx, register_rx) = oneshot::channel::<proto::RegisterConnectorResponse>();
let entry = match self
.socket
.admit(self.key.clone(), inbound_tx, register_tx)
.await
{
Ok(e) => e,
Err(e) => {
logger.error(
&format!(
"ws multiplex handle {}: admission rejected by socket",
self.key
),
&e.to_string(),
);
return Err(e);
}
};
let register_msg = {
let cfg = self.config.read().await;
let token = self.session_token.read().await.clone().unwrap_or_default();
build_register_message_with_token(&cfg, self.connector.as_ref(), &token)
};
if self.socket.outbound_tx.send(register_msg).await.is_err() {
logger.warn(&format!(
"ws multiplex handle {}: outbound closed before register_request could be sent",
self.key
));
self.socket.evict(&self.key, current_arn.as_deref()).await;
return self.maybe_backoff(&logger, &mut attempt).await;
}
let response_deadline =
Duration::from_millis(self.opts.connect_timeout_ms.max(1)).saturating_mul(3);
let resp = match wait_for_register_response(
register_rx,
&self.shutdown,
&self.socket,
response_deadline,
)
.await
{
Ok(r) => r,
Err(e) => {
logger.warn(&format!(
"ws multiplex handle {}: register failed: {e}",
self.key
));
self.socket.evict(&self.key, current_arn.as_deref()).await;
if self.shutdown.load(Ordering::SeqCst) {
return Ok(());
}
if !self.opts.reconnect_enabled {
return Ok(());
}
attempt = attempt.saturating_add(1);
let backoff = compute_backoff(&self.opts, attempt);
{
let mut m = self.metrics.lock().await;
m.reconnection_attempts += 1;
m.current_backoff_ms = backoff.as_millis() as u64;
}
if !sleep_with_shutdown(backoff, &self.shutdown, Some(&self.socket.shutdown))
.await
{
return Ok(());
}
continue;
}
};
if !resp.success {
logger.warn(&format!(
"ws multiplex handle {}: register rejected: status='{}' error='{}'",
self.key, resp.status, resp.error
));
self.socket.evict(&self.key, current_arn.as_deref()).await;
if !self.opts.reconnect_enabled {
return Ok(());
}
attempt = attempt.saturating_add(1);
let backoff = compute_backoff(&self.opts, attempt);
{
let mut m = self.metrics.lock().await;
m.reconnection_attempts += 1;
m.current_backoff_ms = backoff.as_millis() as u64;
}
if !sleep_with_shutdown(backoff, &self.shutdown, Some(&self.socket.shutdown)).await
{
return Ok(());
}
continue;
}
current_arn = Some(resp.connector_arn.clone());
if !resp.session_token.is_empty() {
*self.session_token.write().await = Some(resp.session_token.clone());
}
{
let mut m = self.metrics.lock().await;
m.last_connected_at_ms = Some(now_ms() as u64);
m.current_backoff_ms = 0;
if ever_registered {
m.successful_reconnects += 1;
}
}
ever_registered = true;
self.socket.register_acked.store(true, Ordering::SeqCst);
attempt = 0;
logger.info(&format!(
"ws multiplex handle {}: registered (arn={})",
self.key, resp.connector_arn
));
self.drive_inbound(inbound_rx, &entry, &logger).await;
self.socket.evict(&self.key, current_arn.as_deref()).await;
current_arn = None;
if self.shutdown.load(Ordering::SeqCst) {
return Ok(());
}
if self.socket.shutdown.load(Ordering::SeqCst) {
return Ok(());
}
if !self.opts.reconnect_enabled {
return Ok(());
}
attempt = attempt.saturating_add(1);
let backoff = compute_backoff(&self.opts, attempt);
{
let mut m = self.metrics.lock().await;
m.total_disconnects += 1;
m.last_disconnected_at_ms = Some(now_ms() as u64);
m.last_disconnect_reason = Some("ws-stream-ended".to_string());
m.reconnection_attempts += 1;
m.current_backoff_ms = backoff.as_millis() as u64;
}
if !sleep_with_shutdown(backoff, &self.shutdown, Some(&self.socket.shutdown)).await {
return Ok(());
}
}
}
async fn maybe_backoff(&self, logger: &Logger, attempt: &mut u32) -> Result<()> {
if !self.opts.reconnect_enabled {
return Ok(());
}
*attempt = attempt.saturating_add(1);
let backoff = compute_backoff(&self.opts, *attempt);
logger.warn(&format!(
"ws multiplex handle {}: outbound channel lost; retrying in {}ms (attempt {})",
self.key,
backoff.as_millis(),
attempt
));
if !sleep_with_shutdown(backoff, &self.shutdown, Some(&self.socket.shutdown)).await {
return Ok(());
}
Ok(())
}
async fn drive_inbound(
&self,
mut inbound_rx: mpsc::Receiver<StreamMessage>,
entry: &Arc<HandleEntry>,
logger: &Logger,
) {
loop {
if self.shutdown.load(Ordering::SeqCst) || self.socket.shutdown.load(Ordering::SeqCst) {
return;
}
tokio::select! {
msg_opt = inbound_rx.recv() => {
match msg_opt {
Some(msg) => {
self.dispatch_inbound(msg, entry, logger).await;
}
None => {
return;
}
}
}
_ = tokio::time::sleep(SHUTDOWN_POLL) => {}
}
}
}
async fn dispatch_inbound(
&self,
msg: StreamMessage,
entry: &Arc<HandleEntry>,
logger: &Logger,
) {
match msg.message {
Some(stream_message::Message::ExecuteRequest(req)) => {
let request = SdkExecuteRequest {
request_id: req.request_id.clone(),
payload: req.payload,
payload_encoding: PayloadEncoding::from(req.payload_encoding),
context: req.context,
capability_id: if req.capability_id.is_empty() {
None
} else {
Some(req.capability_id)
},
};
let connector = self.connector.clone();
let metrics = self.metrics.clone();
let outbound = self.socket.outbound_tx.clone();
let key = self.key.clone();
let semaphore = self.request_semaphore.clone();
let exec_logger = Logger::new("multi/ws/handle/execute");
tokio::spawn(async move {
let permit = match semaphore.acquire_owned().await {
Ok(p) => p,
Err(_) => {
exec_logger.debug(&format!(
"ws multiplex handle {key}: request semaphore closed"
));
return;
}
};
if let Err(e) =
handle_execute(connector, request, outbound, metrics, &exec_logger, &key)
.await
{
exec_logger.error(
&format!("ws multiplex handle {key}: execute dispatch failed"),
&e.to_string(),
);
}
drop(permit);
});
}
Some(stream_message::Message::RegisterResponse(resp)) => {
if resp.success {
if !resp.session_token.is_empty() {
*self.session_token.write().await = Some(resp.session_token.clone());
}
if !resp.connector_arn.is_empty() {
self.socket
.bind_arn(resp.connector_arn.clone(), entry.clone())
.await;
}
logger.info(&format!(
"ws multiplex handle {}: in-stream re-register succeeded (arn={})",
self.key, resp.connector_arn
));
} else {
logger.warn(&format!(
"ws multiplex handle {}: in-stream re-register failed: status='{}' error='{}'",
self.key, resp.status, resp.error
));
}
}
Some(stream_message::Message::CredentialsIssued(creds)) => {
let key = self.key.clone();
let config = self.config.clone();
let connector = self.connector.clone();
let socket = self.socket.clone();
let session_token = self.session_token.clone();
let connect_timeout_ms = self.opts.connect_timeout_ms;
let creds_logger = Logger::new("multi/ws/handle/creds");
tokio::spawn(async move {
handle_credentials_issued(
key,
config,
connector,
socket,
session_token,
connect_timeout_ms,
creds,
creds_logger,
)
.await;
});
}
Some(stream_message::Message::ApprovalNotification(notif)) => {
let status = proto::RegistrationStatus::try_from(notif.status);
match status {
Ok(proto::RegistrationStatus::Approved) => {
logger.info(&format!(
"ws multiplex handle {}: approved (CredentialsIssued imminent)",
self.key
));
}
Ok(proto::RegistrationStatus::Pending) => {
logger.info(&format!(
"ws multiplex handle {}: pending approval — {}",
self.key,
if notif.message.is_empty() {
"awaiting admin"
} else {
¬if.message
}
));
}
Ok(proto::RegistrationStatus::Rejected) => {
logger.warn(&format!(
"ws multiplex handle {}: REJECTED — {}",
self.key,
if notif.message.is_empty() {
"no reason"
} else {
¬if.message
}
));
}
_ => {}
}
}
Some(stream_message::Message::HeartbeatRequest(_)) => {
let resp = StreamMessage {
message: Some(stream_message::Message::HeartbeatResponse(
HeartbeatResponse {
gateway_id: String::new(),
timestamp_ms: now_ms(),
should_reconnect: false,
},
)),
};
let _ = self.socket.outbound_tx.send(resp).await;
}
other => {
logger.debug(&format!(
"ws multiplex handle {}: ignoring inbound variant {:?}",
self.key,
other.as_ref().map(std::mem::discriminant)
));
}
}
}
}
#[allow(clippy::too_many_arguments)]
async fn handle_credentials_issued(
key: RegistrationKey,
config: Arc<RwLock<ConnectorConfig>>,
connector: Arc<dyn BaseConnector>,
socket: Arc<WsMultiplexSocket>,
session_token: Arc<RwLock<Option<String>>>,
connect_timeout_ms: u64,
creds: proto::CredentialsIssued,
logger: Logger,
) {
if creds.ott.is_empty() {
logger.warn(&format!(
"ws multiplex handle {key}: CredentialsIssued without OTT",
));
return;
}
if creds.matrix_api_url.is_empty() {
logger.warn(&format!(
"ws multiplex handle {key}: CredentialsIssued without matrix_api_url",
));
return;
}
let http_budget =
Duration::from_millis(connect_timeout_ms.saturating_mul(3).max(connect_timeout_ms));
let (instance_id, connector_type) = (
config.read().await.instance_id.clone(),
connector.connector_type().to_string(),
);
let mut provider =
crate::auth::OttProvider::new(Some(connector_type.clone()), Some(instance_id.clone()));
let register_fut = provider.register_public_key_with_ott_data(
&creds.ott,
&creds.matrix_api_url,
&creds.register_url,
&connector_type,
Some(&instance_id),
);
match tokio::time::timeout(http_budget, register_fut).await {
Ok(Ok(_creds)) => {}
Ok(Err(e)) => {
logger.error(
&format!("ws multiplex handle {key}: OTT public-key registration failed"),
&e.to_string(),
);
return;
}
Err(_) => {
logger.warn(&format!(
"ws multiplex handle {key}: OTT public-key registration timed out after {}ms",
http_budget.as_millis()
));
return;
}
}
let jwt = match tokio::time::timeout(http_budget, provider.get_token()).await {
Ok(Ok(t)) => t,
Ok(Err(e)) => {
logger.error(
&format!("ws multiplex handle {key}: post-OTT JWT exchange failed"),
&e.to_string(),
);
return;
}
Err(_) => {
logger.warn(&format!(
"ws multiplex handle {key}: post-OTT JWT exchange timed out after {}ms",
http_budget.as_millis()
));
return;
}
};
config.write().await.auth_token = jwt;
let cfg = config.read().await.clone();
let token = session_token.read().await.clone().unwrap_or_default();
let msg = build_register_message_with_token(&cfg, connector.as_ref(), &token);
if socket.outbound_tx.send(msg).await.is_err() {
logger.warn(&format!(
"ws multiplex handle {key}: socket closed before in-stream re-register sent",
));
}
}
fn build_register_message_with_token(
config: &ConnectorConfig,
connector: &dyn BaseConnector,
session_token: &str,
) -> StreamMessage {
let mut metadata = crate::connector::build_registration_metadata(connector);
for (k, v) in &config.metadata {
metadata.insert(k.clone(), v.clone());
}
crate::sdk_metadata::merge_into(
&mut metadata,
&config.transport_type.to_string(),
config.use_tls,
);
let capabilities_proto = ConnectorCapabilities {
connector_type: connector.connector_type().to_string(),
version: connector.version().to_string(),
supported_encodings: connector
.supported_encodings()
.iter()
.map(|e| *e as i32)
.collect(),
behaviors: connector.behaviors().iter().map(|b| *b as i32).collect(),
metadata: metadata.clone(),
task_types: {
let caps = connector.capabilities();
if caps.is_empty() {
Vec::new()
} else {
caps.iter()
.map(|tt| proto::TaskTypeSchema {
task_type_id: tt.task_type_id.clone(),
name: tt.name.clone(),
description: tt.description.clone(),
category: tt.category.clone(),
icon: tt.icon.clone(),
input_schema_json: tt.input_schema_json.clone(),
output_schema_json: tt.output_schema_json.clone(),
})
.collect()
}
},
};
let sanitized_instance_id = sanitize_identifier(&config.instance_id);
let instance_metadata = Some(InstanceMetadata {
display_name: config
.display_name
.clone()
.unwrap_or_else(|| sanitized_instance_id.clone()),
tags: config.tags.clone(),
metadata,
});
let register = RegisterConnectorRequest {
tenant_id: sanitize_identifier(&config.tenant_id),
connector_type: sanitize_identifier(connector.connector_type()),
instance_id: sanitized_instance_id,
capabilities: Some(capabilities_proto),
jwt_token: config.auth_token.clone(),
session_token: session_token.to_string(),
scope: 0,
instance_metadata,
};
StreamMessage {
message: Some(stream_message::Message::RegisterRequest(register)),
}
}
async fn handle_execute(
connector: Arc<dyn BaseConnector>,
request: SdkExecuteRequest,
outbound: mpsc::Sender<StreamMessage>,
metrics: Arc<Mutex<ConnectorMetrics>>,
logger: &Logger,
key: &RegistrationKey,
) -> Result<()> {
let start = Instant::now();
{
let mut m = metrics.lock().await;
m.requests_received += 1;
m.bytes_received += request.payload.len() as u64;
m.last_request_at_ms = chrono::Utc::now().timestamp_millis().max(0) as u64;
}
let request_id = request.request_id.clone();
let response = match deserialize_payload::<serde_json::Value>(
&request.payload,
request.payload_encoding,
) {
Ok(payload) => match connector
.execute_with_context(payload, request.capability_id.as_deref(), &request.context)
.await
{
Ok(value) => match serialize_payload(&value, PayloadEncoding::Json) {
Ok(bytes) => proto::ExecuteResponse {
request_id,
success: true,
payload: bytes.clone(),
payload_encoding: PayloadEncoding::Json as i32,
error: String::new(),
duration_ms: start.elapsed().as_millis() as i64,
},
Err(e) => {
logger.error(
&format!("ws multiplex handle {key}: serialize failed"),
&e.to_string(),
);
let mut m = metrics.lock().await;
m.requests_failed += 1;
proto::ExecuteResponse {
request_id,
success: false,
payload: error_response(&e.to_string()).unwrap_or_default(),
payload_encoding: PayloadEncoding::Json as i32,
error: e.to_string(),
duration_ms: start.elapsed().as_millis() as i64,
}
}
},
Err(e) => {
logger.error(
&format!("ws multiplex handle {key}: execute failed"),
&e.to_string(),
);
let mut m = metrics.lock().await;
m.requests_failed += 1;
proto::ExecuteResponse {
request_id,
success: false,
payload: error_response(&e.to_string()).unwrap_or_default(),
payload_encoding: PayloadEncoding::Json as i32,
error: e.to_string(),
duration_ms: start.elapsed().as_millis() as i64,
}
}
},
Err(e) => {
logger.error(
&format!("ws multiplex handle {key}: deserialize failed"),
&e.to_string(),
);
let mut m = metrics.lock().await;
m.requests_failed += 1;
proto::ExecuteResponse {
request_id,
success: false,
payload: error_response(&e.to_string()).unwrap_or_default(),
payload_encoding: PayloadEncoding::Json as i32,
error: e.to_string(),
duration_ms: start.elapsed().as_millis() as i64,
}
}
};
{
let mut m = metrics.lock().await;
if response.success {
m.requests_processed += 1;
m.bytes_sent += response.payload.len() as u64;
}
m.total_duration_ms += response.duration_ms.max(0) as u64;
}
let msg = StreamMessage {
message: Some(stream_message::Message::ExecuteResponse(response)),
};
if outbound.send(msg).await.is_err() {
return Err(ConnectorError::StreamError(
"ws multiplex outbound closed".to_string(),
));
}
Ok(())
}
async fn wait_for_register_response(
register_rx: oneshot::Receiver<proto::RegisterConnectorResponse>,
shutdown: &Arc<AtomicBool>,
socket: &Arc<WsMultiplexSocket>,
deadline: Duration,
) -> Result<proto::RegisterConnectorResponse> {
let mut register_rx = register_rx;
let start = Instant::now();
loop {
if shutdown.load(Ordering::SeqCst) {
return Err(ConnectorError::StreamError(
"shutdown while awaiting register response".to_string(),
));
}
if socket.shutdown.load(Ordering::SeqCst) {
return Err(ConnectorError::StreamError(
"socket closed while awaiting register response".to_string(),
));
}
if start.elapsed() >= deadline {
return Err(ConnectorError::Timeout(
"register response did not arrive in time".to_string(),
));
}
let remaining = deadline.saturating_sub(start.elapsed()).min(SHUTDOWN_POLL);
match tokio::time::timeout(remaining, &mut register_rx).await {
Ok(Ok(resp)) => return Ok(resp),
Ok(Err(_)) => {
return Err(ConnectorError::StreamError(
"register response sender dropped".to_string(),
));
}
Err(_) => continue,
}
}
}
fn compute_backoff(opts: &MultiTransportOptions, attempt: u32) -> Duration {
let base = opts.reconnect_delay_ms;
let max = opts.max_backoff_delay_ms;
let exp = (attempt.saturating_sub(1)).min(20);
let scaled = base.saturating_mul(1u64 << exp);
let jitter = if opts.reconnect_jitter_ms > 0 {
thread_rng().gen_range(0..=opts.reconnect_jitter_ms)
} else {
0
};
let with_jitter = scaled.saturating_add(jitter);
Duration::from_millis(with_jitter.min(max))
}
async fn sleep_with_shutdown(
total: Duration,
shutdown: &Arc<AtomicBool>,
socket_shutdown: Option<&Arc<AtomicBool>>,
) -> bool {
let deadline = Instant::now() + total;
loop {
if shutdown.load(Ordering::SeqCst) {
return false;
}
if let Some(s) = socket_shutdown
&& s.load(Ordering::SeqCst)
{
return false;
}
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
return true;
}
let step = remaining.min(SHUTDOWN_POLL);
tokio::time::sleep(step).await;
}
}
pub(crate) struct WsMultiplexDriver {
pub opts: MultiTransportOptions,
pub shutdown: Arc<AtomicBool>,
}
pub(crate) struct WsMultiplexEntry {
pub key: RegistrationKey,
pub config: ConnectorConfig,
pub connector: Arc<dyn BaseConnector>,
pub metrics: Arc<Mutex<ConnectorMetrics>>,
}
impl WsMultiplexDriver {
pub(crate) async fn run(self, entries: Vec<WsMultiplexEntry>) -> Result<()> {
let logger = Logger::new("multi/ws/driver");
let mut groups: HashMap<String, Vec<WsMultiplexEntry>> = HashMap::new();
for entry in entries {
let sanitized = sanitize_identifier(&entry.config.tenant_id);
groups.entry(sanitized).or_default().push(entry);
}
let mut tenant_tasks = Vec::with_capacity(groups.len());
for (tenant_sanitized, group) in groups {
let opts = self.opts.clone();
let shutdown = self.shutdown.clone();
let logger = Logger::new("multi/ws/driver");
tenant_tasks.push(tokio::spawn(async move {
run_tenant_group(tenant_sanitized, group, opts, shutdown, logger).await
}));
}
for t in tenant_tasks {
match t.await {
Ok(Ok(())) => {}
Ok(Err(e)) => logger.warn(&format!("tenant group exited with error: {e}")),
Err(e) => logger.error("tenant group task panicked", &e.to_string()),
}
}
Ok(())
}
}
async fn run_tenant_group(
tenant_sanitized: String,
entries: Vec<WsMultiplexEntry>,
opts: MultiTransportOptions,
shutdown: Arc<AtomicBool>,
logger: Logger,
) -> Result<()> {
let mut socket_attempt: u32 = 0;
loop {
if shutdown.load(Ordering::SeqCst) {
return Ok(());
}
let (socket, demux_handle) =
match open_socket(&tenant_sanitized, &opts, shutdown.clone()).await {
Ok(pair) => pair,
Err(e) => {
logger.warn(&format!(
"ws multiplex tenant '{tenant_sanitized}': open_socket failed: {e}"
));
if !opts.reconnect_enabled {
return Err(e);
}
socket_attempt = socket_attempt.saturating_add(1);
let backoff = compute_backoff(&opts, socket_attempt);
if !sleep_with_shutdown(backoff, &shutdown, None).await {
return Ok(());
}
continue;
}
};
let mut handle_tasks = Vec::with_capacity(entries.len());
for entry in &entries {
let opts = opts.clone();
let shutdown = shutdown.clone();
let socket = socket.clone();
let key = entry.key.clone();
let metrics = entry.metrics.clone();
let connector = entry.connector.clone();
let mut config = entry.config.clone();
config.transport_type = TransportType::WebSocket;
config.host = opts.host.clone();
config.use_tls = opts.use_tls;
let handle = WsMultiplexHandle {
key,
config: Arc::new(RwLock::new(config)),
connector,
socket,
shutdown,
metrics,
opts: opts.clone(),
request_semaphore: Arc::new(Semaphore::new(opts.max_concurrent_requests.max(1))),
session_token: Arc::new(RwLock::new(None)),
};
handle_tasks.push(tokio::spawn(async move { handle.run().await }));
}
for t in handle_tasks {
match t.await {
Ok(Ok(())) => {}
Ok(Err(e)) => logger.warn(&format!(
"ws multiplex handle exited with error (tenant '{tenant_sanitized}'): {e}"
)),
Err(e) => logger.error(
&format!("ws multiplex handle task panicked (tenant '{tenant_sanitized}')"),
&e.to_string(),
),
}
}
socket.shutdown.store(true, Ordering::SeqCst);
demux_handle.abort();
let _ = demux_handle.await;
if shutdown.load(Ordering::SeqCst) || !opts.reconnect_enabled {
return Ok(());
}
if socket.register_acked.load(Ordering::SeqCst) {
socket_attempt = 0;
}
socket_attempt = socket_attempt.saturating_add(1);
let backoff = compute_backoff(&opts, socket_attempt);
logger.warn(&format!(
"ws multiplex tenant '{tenant_sanitized}': all handles ended; reopening socket in {}ms",
backoff.as_millis()
));
if !sleep_with_shutdown(backoff, &shutdown, None).await {
return Ok(());
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn arn_parts_round_trip() {
assert_eq!(
arn_parts("matrix:demo-org:echo:echo-1"),
Some(("demo-org", "echo", "echo-1"))
);
assert_eq!(arn_parts("matrix:t:c:i"), Some(("t", "c", "i")));
assert_eq!(
arn_parts("matrix:t:c:weird:instance"),
Some(("t", "c", "weird:instance"))
);
}
#[test]
fn arn_parts_rejects_non_matrix() {
assert!(arn_parts("foo:t:c:i").is_none());
assert!(arn_parts("matrix:t:c").is_none());
assert!(arn_parts("").is_none());
}
fn key(instance: &str) -> RegistrationKey {
RegistrationKey {
tenant_id: "t".into(),
connector_type: "c".into(),
instance_id: instance.into(),
}
}
#[tokio::test]
async fn admit_respects_configured_cap() {
let (tx, _rx) = mpsc::channel(1);
let socket = Arc::new(WsMultiplexSocket::new("t".into(), tx, 3));
for i in 0..3 {
let (in_tx, _in_rx) = mpsc::channel(1);
let (reg_tx, _reg_rx) = oneshot::channel();
socket
.admit(key(&format!("inst-{i}")), in_tx, reg_tx)
.await
.expect("admit");
}
let (in_tx, _in_rx) = mpsc::channel(1);
let (reg_tx, _reg_rx) = oneshot::channel();
let err = socket
.admit(key("inst-3"), in_tx, reg_tx)
.await
.map(|_| ())
.expect_err("4th admit should be rejected");
let msg = err.to_string();
assert!(
msg.contains("handler cap of 3"),
"error should name the configured cap, got: {msg}"
);
}
#[tokio::test]
async fn admit_rejects_duplicate_keys() {
let (tx, _rx) = mpsc::channel(1);
let socket = Arc::new(WsMultiplexSocket::new("t".into(), tx, 8));
let (in_tx, _in_rx) = mpsc::channel(1);
let (reg_tx, _reg_rx) = oneshot::channel();
socket
.admit(key("inst-0"), in_tx, reg_tx)
.await
.expect("first admit");
let (in_tx2, _in_rx2) = mpsc::channel(1);
let (reg_tx2, _reg_rx2) = oneshot::channel();
let err = socket
.admit(key("inst-0"), in_tx2, reg_tx2)
.await
.map(|_| ())
.expect_err("duplicate should be rejected");
assert!(err.to_string().contains("already admitted"));
}
#[test]
fn compute_backoff_caps_at_max() {
let opts = MultiTransportOptions {
reconnect_delay_ms: 100,
max_backoff_delay_ms: 5_000,
reconnect_jitter_ms: 0,
..MultiTransportOptions::default()
};
let backoff = compute_backoff(&opts, 21);
assert_eq!(backoff, Duration::from_millis(5_000));
}
async fn admit_with_rx(
socket: &Arc<WsMultiplexSocket>,
instance: &str,
) -> (Arc<HandleEntry>, mpsc::Receiver<StreamMessage>) {
let (in_tx, in_rx) = mpsc::channel(8);
let (reg_tx, _reg_rx) = oneshot::channel();
let entry = socket
.admit(key(instance), in_tx, reg_tx)
.await
.expect("admit");
(entry, in_rx)
}
#[tokio::test]
async fn credentials_issued_routes_by_gateway_id() {
let (out_tx, _out_rx) = mpsc::channel(8);
let socket = Arc::new(WsMultiplexSocket::new("t".into(), out_tx, 8));
let (_entry_a, mut rx_a) = admit_with_rx(&socket, "inst-a").await;
let (_entry_b, mut rx_b) = admit_with_rx(&socket, "inst-b").await;
let logger = Logger::new("test");
let creds = StreamMessage {
message: Some(stream_message::Message::CredentialsIssued(
proto::CredentialsIssued {
gateway_id: "t.c.inst-a".to_string(),
..Default::default()
},
)),
};
route_inbound_by_gateway_id(&socket, creds, "t.c.inst-a", "credentials_issued", &logger)
.await;
assert!(rx_a.try_recv().is_ok(), "addressed handle must receive");
assert!(
rx_b.try_recv().is_err(),
"other handle MUST NOT receive credentials addressed elsewhere"
);
}
#[tokio::test]
async fn credentials_issued_unknown_gateway_drops() {
let (out_tx, _out_rx) = mpsc::channel(8);
let socket = Arc::new(WsMultiplexSocket::new("t".into(), out_tx, 8));
let (_entry, mut rx) = admit_with_rx(&socket, "inst-survivor").await;
let logger = Logger::new("test");
let creds = StreamMessage {
message: Some(stream_message::Message::CredentialsIssued(
proto::CredentialsIssued {
gateway_id: "t.c.inst-evicted".to_string(),
..Default::default()
},
)),
};
route_inbound_by_gateway_id(
&socket,
creds,
"t.c.inst-evicted",
"credentials_issued",
&logger,
)
.await;
assert!(
rx.try_recv().is_err(),
"credentials for unknown gateway must be dropped, not delivered"
);
}
#[tokio::test]
async fn unmatched_arn_does_not_fall_back_to_sole_survivor() {
let (out_tx, _out_rx) = mpsc::channel(8);
let socket = Arc::new(WsMultiplexSocket::new("t".into(), out_tx, 8));
let (_entry_survivor, mut rx_survivor) = admit_with_rx(&socket, "inst-survivor").await;
let logger = Logger::new("test");
let exec_msg = StreamMessage {
message: Some(stream_message::Message::ExecuteRequest(
proto::ExecuteRequest {
request_id: "req-1".into(),
..Default::default()
},
)),
};
route_inbound_by_arn(
&socket,
exec_msg,
Some("matrix:t:c:inst-evicted"),
"execute_request",
&logger,
)
.await;
assert!(
rx_survivor.try_recv().is_err(),
"addressed-but-unmatched ARN MUST be dropped, not delivered to survivor"
);
}
#[tokio::test]
async fn missing_arn_with_sole_handle_falls_back() {
let (out_tx, _out_rx) = mpsc::channel(8);
let socket = Arc::new(WsMultiplexSocket::new("t".into(), out_tx, 8));
let (_entry, mut rx) = admit_with_rx(&socket, "inst-only").await;
let logger = Logger::new("test");
let exec_msg = StreamMessage {
message: Some(stream_message::Message::ExecuteRequest(
proto::ExecuteRequest::default(),
)),
};
route_inbound_by_arn(&socket, exec_msg, None, "execute_request", &logger).await;
assert!(
rx.try_recv().is_ok(),
"no-ARN frame must reach sole admitted handle"
);
}
#[tokio::test]
async fn register_response_routes_by_instance_id_when_type_canonicalised() {
let (out_tx, _out_rx) = mpsc::channel(8);
let socket = Arc::new(WsMultiplexSocket::new("t".into(), out_tx, 8));
let (in_tx, _in_rx) = mpsc::channel(8);
let (reg_tx, reg_rx) = oneshot::channel();
let admitted_key = RegistrationKey {
tenant_id: "t".into(),
connector_type: "construct-jira-mcp".into(),
instance_id: "host-jira-mcp".into(),
};
socket
.admit(admitted_key, in_tx, reg_tx)
.await
.expect("admit");
let (in_tx2, _in_rx2) = mpsc::channel(8);
let (reg_tx2, _reg_rx2) = oneshot::channel();
let other_key = RegistrationKey {
tenant_id: "t".into(),
connector_type: "construct-other".into(),
instance_id: "host-other".into(),
};
socket
.admit(other_key, in_tx2, reg_tx2)
.await
.expect("admit");
let logger = Logger::new("test");
let resp_msg = StreamMessage {
message: Some(stream_message::Message::RegisterResponse(
proto::RegisterConnectorResponse {
success: true,
connector_arn: "matrix:t:jira-mcp:host-jira-mcp".into(),
..Default::default()
},
)),
};
dispatch_inbound_to_handle(&socket, resp_msg, &logger).await;
let delivered = reg_rx.await.expect("pending oneshot fulfilled");
assert!(delivered.success);
assert_eq!(delivered.connector_arn, "matrix:t:jira-mcp:host-jira-mcp");
}
#[tokio::test]
async fn register_response_with_ambiguous_instance_id_drops() {
let (out_tx, _out_rx) = mpsc::channel(8);
let socket = Arc::new(WsMultiplexSocket::new("t".into(), out_tx, 8));
let (in_tx_a, _in_rx_a) = mpsc::channel(8);
let (reg_tx_a, mut reg_rx_a) = oneshot::channel();
socket
.admit(
RegistrationKey {
tenant_id: "t".into(),
connector_type: "type-a".into(),
instance_id: "shared-inst".into(),
},
in_tx_a,
reg_tx_a,
)
.await
.expect("admit");
let (in_tx_b, _in_rx_b) = mpsc::channel(8);
let (reg_tx_b, mut reg_rx_b) = oneshot::channel();
socket
.admit(
RegistrationKey {
tenant_id: "t".into(),
connector_type: "type-b".into(),
instance_id: "shared-inst".into(),
},
in_tx_b,
reg_tx_b,
)
.await
.expect("admit");
let logger = Logger::new("test");
let resp_msg = StreamMessage {
message: Some(stream_message::Message::RegisterResponse(
proto::RegisterConnectorResponse {
success: true,
connector_arn: "matrix:t:type-mystery:shared-inst".into(),
..Default::default()
},
)),
};
dispatch_inbound_to_handle(&socket, resp_msg, &logger).await;
assert!(reg_rx_a.try_recv().is_err());
assert!(reg_rx_b.try_recv().is_err());
}
}