use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use std::time::Instant;
use anyhow::anyhow;
use futures::SinkExt;
use futures::StreamExt;
use kcmc::ModelingCmd;
use kcmc::websocket::BatchResponse;
use kcmc::websocket::FailureWebSocketResponse;
use kcmc::websocket::ModelingCmdReq;
use kcmc::websocket::ModelingSessionData;
use kcmc::websocket::OkWebSocketResponseData;
use kcmc::websocket::SuccessWebSocketResponse;
use kcmc::websocket::WebSocketRequest;
use kcmc::websocket::WebSocketResponse;
use kittycad_modeling_cmds as kcmc;
use tokio::sync::RwLock;
use tokio::sync::mpsc;
use tokio::sync::oneshot;
use tokio_tungstenite::tungstenite::Message as WsMsg;
use uuid::Uuid;
use super::EngineTransport;
use super::ResponseInformation;
use super::SocketHealth;
use super::TransportCloseError;
use crate::SourceRange;
use crate::errors::KclError;
use crate::errors::KclErrorDetails;
use crate::log::logln;
pub struct TcpRead {
stream: futures::stream::SplitStream<tokio_tungstenite::WebSocketStream<reqwest::Upgraded>>,
}
impl TcpRead {
pub async fn read(&mut self) -> std::result::Result<WebSocketResponse, WebSocketReadError> {
let Some(msg) = self.stream.next().await else {
return Err(anyhow!("Failed to read from WebSocket").into());
};
let msg = match msg {
Ok(msg) => msg,
Err(e) if matches!(e, tokio_tungstenite::tungstenite::Error::Protocol(_)) => {
return Err(WebSocketReadError::Read(e));
}
Err(e) => return Err(anyhow!("Error reading from engine's WebSocket: {e}").into()),
};
match msg {
WsMsg::Text(text) => kcl_engine_codec::deserialize_response_json(&text)
.map_err(anyhow::Error::from)
.map_err(WebSocketReadError::from),
WsMsg::Binary(bin) => kcl_engine_codec::deserialize_response_msgpack(&bin)
.map_err(anyhow::Error::from)
.map_err(WebSocketReadError::from),
WsMsg::Close(close_frame) => {
let err_msg = close_frame
.map(|frame| frame.reason.to_string())
.unwrap_or("WebSocket closed without specifying a reason.".to_string());
Err(anyhow!(err_msg).into())
}
other => Err(anyhow!("Unexpected WebSocket message from engine API: {other}").into()),
}
}
}
type WebSocketTcpWrite = futures::stream::SplitSink<tokio_tungstenite::WebSocketStream<reqwest::Upgraded>, WsMsg>;
const CLOSE_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(1);
pub struct WebSocketTransport {
tcp_read_handle: tokio::task::AbortHandle,
tcp_write_handle: tokio::task::AbortHandle,
engine_req_tx: mpsc::Sender<ToEngineReq>,
shutdown_tx: mpsc::Sender<()>,
responses: ResponseInformation,
pending_errors: Arc<RwLock<Vec<String>>>,
session_data: Arc<RwLock<Option<ModelingSessionData>>>,
socket_health: Arc<RwLock<SocketHealth>>,
upgrade_request_id: Option<String>,
}
impl Drop for WebSocketTransport {
fn drop(&mut self) {
self.tcp_read_handle.abort();
self.tcp_write_handle.abort();
}
}
#[allow(clippy::large_enum_variant)]
pub enum WebSocketReadError {
Read(tokio_tungstenite::tungstenite::Error),
Deser(anyhow::Error),
}
impl From<anyhow::Error> for WebSocketReadError {
fn from(e: anyhow::Error) -> Self {
Self::Deser(e)
}
}
struct ToEngineReq {
req: WebSocketRequest,
request_sent: oneshot::Sender<anyhow::Result<()>>,
}
impl WebSocketTransport {
pub async fn spawn(
ws: reqwest::Upgraded,
heartbeats: Option<u64>,
response_information: ResponseInformation,
session_data: Arc<RwLock<Option<ModelingSessionData>>>,
pending_errors: Arc<RwLock<Vec<String>>>,
socket_health: Arc<RwLock<SocketHealth>>,
upgrade_request_id: Option<String>,
) -> Self {
let wsconfig = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default()
.max_message_size(Some(usize::MAX))
.max_frame_size(Some(usize::MAX));
let ws_stream = tokio_tungstenite::WebSocketStream::from_raw_socket(
ws,
tokio_tungstenite::tungstenite::protocol::Role::Client,
Some(wsconfig),
)
.await;
let (tcp_write, tcp_read) = ws_stream.split();
let (engine_req_tx, engine_req_rx) = mpsc::channel(10);
let (shutdown_tx, shutdown_rx) = mpsc::channel(1);
let tcp_write_handle = tokio::task::spawn(Self::start_write_actor(
tcp_write,
engine_req_rx,
shutdown_rx,
heartbeats,
));
let mut tcp_read = TcpRead { stream: tcp_read };
let response_information_for_read = response_information.clone();
let session_data_for_read = session_data.clone();
let pending_errors_for_read = pending_errors.clone();
let socket_health_tcp_read = socket_health.clone();
let tcp_read_handle: tokio::task::JoinHandle<Result<(), WebSocketReadError>> = tokio::spawn(async move {
loop {
match tcp_read.read().await {
Ok(ws_resp) => {
let id = ws_resp.request_id();
match &ws_resp {
WebSocketResponse::Success(SuccessWebSocketResponse {
resp: OkWebSocketResponseData::ModelingBatch { responses },
..
}) => {
#[expect(
clippy::iter_over_hash_type,
reason = "modeling command uses a HashMap and keys are random, so we don't really have a choice"
)]
for (resp_id, batch_response) in responses {
let id: uuid::Uuid = (*resp_id).into();
match batch_response {
BatchResponse::Success { response } => {
response_information_for_read
.add(
id,
WebSocketResponse::Success(SuccessWebSocketResponse {
success: true,
request_id: Some(id),
resp: OkWebSocketResponseData::Modeling {
modeling_response: response.clone(),
},
}),
)
.await;
}
BatchResponse::Failure { errors } => {
response_information_for_read
.add(
id,
WebSocketResponse::Failure(FailureWebSocketResponse {
success: false,
request_id: Some(id),
errors: errors.clone(),
}),
)
.await;
}
}
}
}
WebSocketResponse::Success(SuccessWebSocketResponse {
resp: OkWebSocketResponseData::ModelingSessionData { session },
..
}) => {
let mut sd = session_data_for_read.write().await;
sd.replace(session.clone());
logln!("API Call ID: {}", session.api_call_id);
}
WebSocketResponse::Failure(FailureWebSocketResponse {
success: _,
request_id,
errors,
}) => {
if let Some(id) = request_id {
response_information_for_read
.add(
*id,
WebSocketResponse::Failure(FailureWebSocketResponse {
success: false,
request_id: *request_id,
errors: errors.clone(),
}),
)
.await;
} else {
let mut pe = pending_errors_for_read.write().await;
for error in errors {
if !pe.contains(&error.message) {
pe.push(error.message.clone());
}
}
drop(pe);
}
}
WebSocketResponse::Success(SuccessWebSocketResponse {
resp: _debug @ OkWebSocketResponseData::Debug { .. },
..
}) => {
}
_ => {}
}
if let Some(id) = id {
response_information_for_read.add(id, ws_resp.clone()).await;
}
}
Err(e) => {
let msg = match &e {
WebSocketReadError::Read(e) => e.to_string(),
WebSocketReadError::Deser(e) => e.to_string(),
};
pending_errors_for_read.write().await.push(msg);
*socket_health_tcp_read.write().await = SocketHealth::Inactive;
return Err(e);
}
}
}
});
Self {
shutdown_tx,
responses: response_information,
pending_errors,
session_data,
socket_health,
upgrade_request_id,
engine_req_tx,
tcp_read_handle: tcp_read_handle.abort_handle(),
tcp_write_handle: tcp_write_handle.abort_handle(),
}
}
fn connection_id_message(&self, session: Option<&ModelingSessionData>) -> String {
if let Some(session) = session {
format!(" (API call ID: {})", session.api_call_id)
} else if let Some(id) = &self.upgrade_request_id {
format!(" (Engine upgrade request ID: {id})")
} else {
" (No API call ID: session data empty)".to_string()
}
}
async fn inner_send_to_engine_binary(
request: WebSocketRequest,
tcp_write: &mut WebSocketTcpWrite,
) -> anyhow::Result<()> {
let msg = kcl_engine_codec::serialize_request_msgpack(&request)
.map_err(|e| anyhow!("could not serialize msgpack: {e}"))?;
tcp_write
.send(WsMsg::Binary(msg.into()))
.await
.map_err(|e| anyhow!("could not send MsgPack over websocket: {e}"))?;
Ok(())
}
async fn inner_send_to_engine(request: WebSocketRequest, tcp_write: &mut WebSocketTcpWrite) -> anyhow::Result<()> {
let msg =
kcl_engine_codec::serialize_request_json(&request).map_err(|e| anyhow!("could not serialize json: {e}"))?;
tcp_write
.send(WsMsg::Text(msg.into()))
.await
.map_err(|e| anyhow!("could not send json over websocket: {e}"))?;
Ok(())
}
async fn start_write_actor(
mut tcp_write: WebSocketTcpWrite,
mut engine_req_rx: mpsc::Receiver<ToEngineReq>,
mut shutdown_rx: mpsc::Receiver<()>,
heartbeats: Option<u64>,
) {
let heartbeats = heartbeats.unwrap_or_default();
let send_heartbeats = heartbeats != 0;
let period_seconds = if heartbeats == 0 { 5 * 60 } else { heartbeats };
let period = Duration::from_secs(period_seconds);
let mut heartbeats_stream = tokio::time::interval(period);
loop {
tokio::select! {
maybe_req = engine_req_rx.recv() => {
match maybe_req {
Some(ToEngineReq { req, request_sent }) => {
let res = if matches!(
&req,
WebSocketRequest::ModelingCmdReq(ModelingCmdReq {
cmd: ModelingCmd::ImportFiles { .. },
cmd_id: _,
})
) {
Self::inner_send_to_engine_binary(req, &mut tcp_write).await
} else {
Self::inner_send_to_engine(req, &mut tcp_write).await
};
let _ = request_sent.send(res);
}
None => {
break;
}
}
},
_ = shutdown_rx.recv() => {
let _ = Self::inner_close_engine(&mut tcp_write).await;
return;
}
_ = heartbeats_stream.tick(), if send_heartbeats => {
let res = Self::inner_send_to_engine(WebSocketRequest::Ping {}, &mut tcp_write).await;
let _ = res;
}
}
}
let _ = Self::inner_close_engine(&mut tcp_write).await;
}
async fn inner_close_engine(tcp_write: &mut WebSocketTcpWrite) -> anyhow::Result<()> {
tcp_write
.send(WsMsg::Close(None))
.await
.map_err(|e| anyhow!("could not send close over websocket: {e}"))?;
Ok(())
}
async fn close_with_timeout(&self, timeout: Duration) {
let _ = self.shutdown_tx.try_send(());
let _ = tokio::time::timeout(timeout, async {
loop {
if *self.socket_health.read().await == SocketHealth::Inactive {
return;
}
tokio::task::yield_now().await;
}
})
.await;
self.tcp_read_handle.abort();
self.tcp_write_handle.abort();
*self.socket_health.write().await = SocketHealth::Inactive;
}
}
#[async_trait::async_trait]
impl EngineTransport for WebSocketTransport {
async fn inner_fire_modeling_cmd(
&self,
_cmd_id: uuid::Uuid,
source_range: SourceRange,
cmd: WebSocketRequest,
_id_to_source_range: HashMap<Uuid, SourceRange>,
) -> Result<(), KclError> {
let (tx, rx) = oneshot::channel();
let api_call_id_msg = self.connection_id_message(self.session_data.read().await.as_ref());
self.engine_req_tx
.send(ToEngineReq {
req: cmd.clone(),
request_sent: tx,
})
.await
.map_err(|e| {
KclError::new_engine(KclErrorDetails::new(
format!("Failed to send modeling command: {e}{api_call_id_msg}"),
vec![source_range],
))
})?;
let send_result = rx.await.map_err(|e| {
KclError::new_engine_hangup(
KclErrorDetails::new(
format!("could not send request to the engine actor: {e}{api_call_id_msg}"),
vec![source_range],
),
None,
)
})?;
if let Err(send_error) = send_result {
let pending_errors = self.pending_errors.read().await;
if !pending_errors.is_empty() {
return Err(KclError::new_engine(KclErrorDetails::new(
format!("{}{}", pending_errors.join(", "), api_call_id_msg),
vec![source_range],
)));
}
return Err(KclError::new_engine_hangup(
KclErrorDetails::new(
format!("could not send request to the engine: {send_error}{api_call_id_msg}"),
vec![source_range],
),
None,
));
}
Ok(())
}
async fn inner_send_modeling_cmd(
&self,
cmd_id: uuid::Uuid,
source_range: SourceRange,
cmd: WebSocketRequest,
id_to_source_range: HashMap<Uuid, SourceRange>,
) -> Result<WebSocketResponse, KclError> {
self.inner_fire_modeling_cmd(cmd_id, source_range, cmd, id_to_source_range)
.await?;
let response_timeout = 600;
let current_time = Instant::now();
while current_time.elapsed().as_secs() < response_timeout {
let guard = self.socket_health.read().await;
if *guard == SocketHealth::Inactive {
let session_data = self.session_data.read().await;
let api_call_id = session_data.as_ref().map(|session| session.api_call_id.to_string());
let api_call_id_msg = self.connection_id_message(session_data.as_ref());
let pe = self.pending_errors.read().await;
if !pe.is_empty() {
return Err(KclError::new_engine(KclErrorDetails::new(
format!("{}{}", pe.join(", "), api_call_id_msg),
vec![source_range],
)));
} else {
return Err(KclError::new_engine_hangup(
KclErrorDetails::new(
format!("Modeling command failed: websocket closed early{}", api_call_id_msg),
vec![source_range],
),
api_call_id,
));
}
}
if let Some(resp) = self.responses.responses.read().await.get(&cmd_id) {
return Ok(resp.clone());
}
}
let api_call_id_msg = self.connection_id_message(self.session_data.read().await.as_ref());
Err(KclError::new_engine(KclErrorDetails::new(
format!("Modeling command timed out `{cmd_id}`{}", api_call_id_msg),
vec![source_range],
)))
}
async fn close(&self) -> Result<(), TransportCloseError> {
self.close_with_timeout(CLOSE_HANDSHAKE_TIMEOUT).await;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn close_aborts_writer_when_reader_stops_before_timeout() {
let write = tokio::spawn(std::future::pending::<()>());
let (engine_req_tx, _engine_req_rx) = mpsc::channel(1);
let (shutdown_tx, mut shutdown_rx) = mpsc::channel(1);
let socket_health = Arc::new(RwLock::new(SocketHealth::Active));
let read_health = socket_health.clone();
let read = tokio::spawn(async move {
shutdown_rx.recv().await.expect("shutdown should be requested");
*read_health.write().await = SocketHealth::Inactive;
});
let transport = WebSocketTransport {
tcp_read_handle: read.abort_handle(),
tcp_write_handle: write.abort_handle(),
engine_req_tx,
shutdown_tx,
responses: ResponseInformation::new(Arc::new(RwLock::new(Default::default()))),
pending_errors: Arc::new(RwLock::new(Vec::new())),
session_data: Arc::new(RwLock::new(None)),
socket_health: socket_health.clone(),
upgrade_request_id: None,
};
tokio::time::timeout(Duration::from_secs(1), async {
transport.close_with_timeout(Duration::from_secs(10)).await;
read.await.expect("reader should finish normally");
assert!(write.await.unwrap_err().is_cancelled());
})
.await
.expect("close should abort the writer even when the reader has stopped");
assert_eq!(*socket_health.read().await, SocketHealth::Inactive);
}
#[tokio::test]
async fn close_aborts_tasks_when_peer_does_not_close() {
let read = tokio::spawn(std::future::pending::<()>());
let write = tokio::spawn(std::future::pending::<()>());
let (engine_req_tx, _engine_req_rx) = mpsc::channel(1);
let (shutdown_tx, _shutdown_rx) = mpsc::channel(1);
let socket_health = Arc::new(RwLock::new(SocketHealth::Active));
let transport = WebSocketTransport {
tcp_read_handle: read.abort_handle(),
tcp_write_handle: write.abort_handle(),
engine_req_tx,
shutdown_tx,
responses: ResponseInformation::new(Arc::new(RwLock::new(Default::default()))),
pending_errors: Arc::new(RwLock::new(Vec::new())),
session_data: Arc::new(RwLock::new(None)),
socket_health: socket_health.clone(),
upgrade_request_id: None,
};
tokio::time::timeout(
Duration::from_secs(1),
transport.close_with_timeout(Duration::from_millis(1)),
)
.await
.expect("close should be bounded");
assert_eq!(*socket_health.read().await, SocketHealth::Inactive);
assert!(read.await.unwrap_err().is_cancelled());
assert!(write.await.unwrap_err().is_cancelled());
}
}