Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 33 additions & 7 deletions rust/src/napi/convert.rs
Original file line number Diff line number Diff line change
Expand Up @@ -538,6 +538,26 @@ fn js_object_to_multipart_body(
}))
}

fn js_object_to_stream_body(
cx: &mut FunctionContext,
obj: Handle<JsObject>,
) -> NeonResult<Option<RequestBody>> {
let Some(value) = obj.get_opt::<JsValue, _, _>(cx, "bodyStream")? else {
return Ok(None);
};
let stream = value.downcast::<JsObject, _>(cx).or_throw(cx)?;
let handle_value: Handle<JsValue> = stream.get(cx, "uploadHandle")?;
let handle = js_value_to_safe_u64(cx, handle_value, "bodyStream.uploadHandle")?;
let length = stream
.get_opt::<JsValue, _, _>(cx, "length")?
.map(|value| js_value_to_safe_u64(cx, value, "bodyStream.length"))
.transpose()?;
let receiver =
take_upload_receiver(handle).or_else(|error| cx.throw_error(error.to_string()))?;

Ok(Some(RequestBody::Stream { receiver, length }))
}

pub(crate) fn js_object_to_request_options(
cx: &mut FunctionContext,
obj: Handle<JsObject>,
Expand Down Expand Up @@ -571,16 +591,22 @@ pub(crate) fn js_object_to_request_options(
.map(|value| js_value_to_bytes(cx, value))
.transpose()?;
let multipart = js_object_to_multipart_body(cx, obj)?;
let stream = js_object_to_stream_body(cx, obj)?;

let body_count = usize::from(body_bytes.is_some())
+ usize::from(multipart.is_some())
+ usize::from(stream.is_some());

if body_bytes.is_some() && multipart.is_some() {
return cx.throw_type_error("body and multipart cannot both be provided");
if body_count > 1 {
return cx.throw_type_error("body, bodyStream, and multipart are mutually exclusive");
}

let body = match (body_bytes, multipart) {
(Some(bytes), None) => Some(RequestBody::Bytes(bytes)),
(None, Some(multipart)) => Some(RequestBody::Multipart(multipart)),
(None, None) => None,
(Some(_), Some(_)) => unreachable!(),
let body = match (body_bytes, multipart, stream) {
(Some(bytes), None, None) => Some(RequestBody::Bytes(bytes)),
(None, Some(multipart), None) => Some(RequestBody::Multipart(multipart)),
(None, None, Some(stream)) => Some(stream),
(None, None, None) => None,
_ => unreachable!(),
};

let proxy = obj
Expand Down
42 changes: 38 additions & 4 deletions rust/src/napi/websocket.rs
Original file line number Diff line number Diff line change
@@ -1,29 +1,55 @@
use crate::napi::convert::{js_object_to_websocket_options, websocket_to_js_object};
use crate::store::runtime::runtime;
use crate::store::websocket_connect_store::{
cancel_websocket_connect, insert_websocket_connect, remove_websocket_connect,
};
use crate::store::websocket_store::{
close_websocket, read_websocket_message, send_websocket_binary, send_websocket_text,
terminate_websocket,
};
use crate::transport::{connect_websocket, types::WebSocketReadResult};
use crate::transport::{make_websocket, types::WebSocketReadResult};
use neon::prelude::*;
use neon::types::buffer::TypedArray;
use neon::types::JsBuffer;

fn websocket_connect_js(mut cx: FunctionContext) -> JsResult<JsPromise> {
fn websocket_connect_js(mut cx: FunctionContext) -> JsResult<JsObject> {
let options_obj = cx.argument::<JsObject>(0)?;
let options = js_object_to_websocket_options(&mut cx, options_obj)?;

let channel = cx.channel();
let (deferred, promise) = cx.promise();
let (cancel_tx, cancel_rx) = tokio::sync::oneshot::channel::<()>();
let handle = insert_websocket_connect(cancel_tx);

std::thread::spawn(move || {
let result = connect_websocket(options);
let result = runtime().block_on(async move {
tokio::select! {
result = make_websocket(options) => result,
_ = cancel_rx => Err(anyhow::anyhow!("WebSocket connection aborted")),
}
});

remove_websocket_connect(handle);

deferred.settle_with(&channel, move |mut cx| match result {
Ok(websocket) => websocket_to_js_object(&mut cx, websocket),
Err(error) => cx.throw_error(format!("{:#}", error)),
});
});

Ok(promise)
let result = JsObject::new(&mut cx);
let handle_value = cx.number(handle as f64);

result.set(&mut cx, "handle", handle_value)?;
result.set(&mut cx, "promise", promise)?;

Ok(result)
}

fn websocket_cancel_connect_js(mut cx: FunctionContext) -> JsResult<JsBoolean> {
let handle = cx.argument::<JsNumber>(0)?.value(&mut cx) as u64;

Ok(cx.boolean(cancel_websocket_connect(handle)))
}

fn websocket_read_js(mut cx: FunctionContext) -> JsResult<JsPromise> {
Expand Down Expand Up @@ -140,11 +166,19 @@ fn websocket_close_js(mut cx: FunctionContext) -> JsResult<JsPromise> {
Ok(promise)
}

fn websocket_terminate_js(mut cx: FunctionContext) -> JsResult<JsBoolean> {
let handle = cx.argument::<JsNumber>(0)?.value(&mut cx) as u64;

Ok(cx.boolean(terminate_websocket(handle)))
}

pub fn register(cx: &mut ModuleContext) -> NeonResult<()> {
cx.export_function("websocketConnect", websocket_connect_js)?;
cx.export_function("websocketCancelConnect", websocket_cancel_connect_js)?;
cx.export_function("websocketRead", websocket_read_js)?;
cx.export_function("websocketSendText", websocket_send_text_js)?;
cx.export_function("websocketSendBinary", websocket_send_binary_js)?;
cx.export_function("websocketClose", websocket_close_js)?;
cx.export_function("websocketTerminate", websocket_terminate_js)?;
Ok(())
}
1 change: 1 addition & 0 deletions rust/src/store/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,5 @@ pub mod client_store;
pub mod request_store;
pub mod runtime;
pub mod upload_store;
pub mod websocket_connect_store;
pub mod websocket_store;
42 changes: 42 additions & 0 deletions rust/src/store/websocket_connect_store.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
use std::collections::HashMap;
use std::sync::{
atomic::{AtomicU64, Ordering},
Mutex, OnceLock,
};

static NEXT_WEBSOCKET_CONNECT_HANDLE: AtomicU64 = AtomicU64::new(1);
static WEBSOCKET_CONNECT_STORE: OnceLock<Mutex<HashMap<u64, tokio::sync::oneshot::Sender<()>>>> =
OnceLock::new();

fn websocket_connect_store() -> &'static Mutex<HashMap<u64, tokio::sync::oneshot::Sender<()>>> {
WEBSOCKET_CONNECT_STORE.get_or_init(|| Mutex::new(HashMap::new()))
}

pub fn insert_websocket_connect(cancel: tokio::sync::oneshot::Sender<()>) -> u64 {
let handle = NEXT_WEBSOCKET_CONNECT_HANDLE.fetch_add(1, Ordering::Relaxed);

websocket_connect_store()
.lock()
.expect("websocket connect store poisoned")
.insert(handle, cancel);

handle
}

pub fn remove_websocket_connect(handle: u64) {
websocket_connect_store()
.lock()
.expect("websocket connect store poisoned")
.remove(&handle);
}

pub fn cancel_websocket_connect(handle: u64) -> bool {
let cancel = websocket_connect_store()
.lock()
.expect("websocket connect store poisoned")
.remove(&handle);

cancel
.map(|cancel| cancel.send(()).is_ok())
.unwrap_or(false)
}
64 changes: 49 additions & 15 deletions rust/src/store/websocket_store.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,17 +8,24 @@ use std::sync::{

#[derive(Debug)]
pub(crate) enum WebSocketCommand {
Text(String),
Binary(Vec<u8>),
Text {
text: String,
ack: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
},
Binary {
bytes: Vec<u8>,
ack: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
},
Close {
code: Option<u16>,
reason: Option<String>,
ack: tokio::sync::oneshot::Sender<std::result::Result<(), String>>,
},
}

#[derive(Debug)]
pub(crate) struct StoredWebSocket {
pub commands: tokio::sync::mpsc::UnboundedSender<WebSocketCommand>,
pub commands: tokio::sync::mpsc::Sender<WebSocketCommand>,
pub events: tokio::sync::Mutex<tokio::sync::mpsc::UnboundedReceiver<WebSocketReadResult>>,
}

Expand All @@ -32,7 +39,7 @@ fn websocket_store() -> &'static Mutex<HashMap<u64, SharedWebSocket>> {
}

pub(crate) fn insert_websocket(
commands: tokio::sync::mpsc::UnboundedSender<WebSocketCommand>,
commands: tokio::sync::mpsc::Sender<WebSocketCommand>,
events: tokio::sync::mpsc::UnboundedReceiver<WebSocketReadResult>,
) -> u64 {
let handle = NEXT_WEBSOCKET_HANDLE.fetch_add(1, Ordering::Relaxed);
Expand Down Expand Up @@ -62,11 +69,16 @@ fn get_websocket(handle: u64) -> Result<SharedWebSocket> {
.ok_or_else(|| anyhow::anyhow!("Unknown websocket handle: {}", handle))
}

pub(crate) fn remove_websocket(handle: u64) {
pub(crate) fn remove_websocket(handle: u64) -> bool {
websocket_store()
.lock()
.expect("websocket store poisoned")
.remove(&handle);
.remove(&handle)
.is_some()
}

pub fn terminate_websocket(handle: u64) -> bool {
remove_websocket(handle)
}

pub fn read_websocket_message(handle: u64) -> Result<WebSocketReadResult> {
Expand All @@ -87,27 +99,49 @@ pub fn read_websocket_message(handle: u64) -> Result<WebSocketReadResult> {
result
}

fn send_websocket_command(handle: u64, command: WebSocketCommand) -> Result<()> {
fn send_websocket_command(
handle: u64,
command: WebSocketCommand,
acknowledgement: tokio::sync::oneshot::Receiver<std::result::Result<(), String>>,
) -> Result<()> {
let websocket = get_websocket(handle)?;
let result = crate::store::runtime::runtime().block_on(async {
websocket
.commands
.send(command)
.await
.map_err(|_| anyhow::anyhow!("WebSocket is already closed"))?;

if websocket.commands.send(command).is_err() {
acknowledgement
.await
.map_err(|_| anyhow::anyhow!("WebSocket send acknowledgement was dropped"))?
.map_err(anyhow::Error::msg)
});

if result.is_err() {
remove_websocket(handle);
anyhow::bail!("WebSocket is already closed");
}

Ok(())
result
}

pub fn send_websocket_text(handle: u64, text: String) -> Result<()> {
send_websocket_command(handle, WebSocketCommand::Text(text))
let (ack, result) = tokio::sync::oneshot::channel();
send_websocket_command(handle, WebSocketCommand::Text { text, ack }, result)
}

pub fn send_websocket_binary(handle: u64, bytes: Vec<u8>) -> Result<()> {
send_websocket_command(handle, WebSocketCommand::Binary(bytes))
let (ack, result) = tokio::sync::oneshot::channel();
send_websocket_command(handle, WebSocketCommand::Binary { bytes, ack }, result)
}

pub fn close_websocket(handle: u64, code: Option<u16>, reason: Option<String>) -> Result<()> {
send_websocket_command(handle, WebSocketCommand::Close { code, reason })
let (ack, result) = tokio::sync::oneshot::channel();
send_websocket_command(
handle,
WebSocketCommand::Close { code, reason, ack },
result,
)
}

#[cfg(test)]
Expand All @@ -116,7 +150,7 @@ mod tests {

#[test]
fn removes_websocket_when_event_stream_closes() {
let (commands, _command_receiver) = tokio::sync::mpsc::unbounded_channel();
let (commands, _command_receiver) = tokio::sync::mpsc::channel(1);
let (event_sender, events) = tokio::sync::mpsc::unbounded_channel();
let handle = insert_websocket(commands, events);

Expand All @@ -133,7 +167,7 @@ mod tests {

#[test]
fn removes_websocket_when_command_stream_closes() {
let (commands, command_receiver) = tokio::sync::mpsc::unbounded_channel();
let (commands, command_receiver) = tokio::sync::mpsc::channel(1);
let (_event_sender, events) = tokio::sync::mpsc::unbounded_channel();
let handle = insert_websocket(commands, events);

Expand Down
2 changes: 1 addition & 1 deletion rust/src/transport/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,4 +8,4 @@ pub mod types;
mod websocket;

pub use request::make_request;
pub use websocket::connect_websocket;
pub(crate) use websocket::make_websocket;
13 changes: 13 additions & 0 deletions rust/src/transport/request.rs
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,19 @@ pub async fn make_request(options: RequestOptions) -> Result<Response> {
if let Some(body) = body {
request = match body {
RequestBody::Bytes(bytes) => request.body(bytes),
RequestBody::Stream { receiver, length } => {
if let Some(length) = length {
let has_content_length = headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("content-length"));

if !has_content_length {
request = request.header("content-length", length);
}
}

request.body(Body::wrap_stream(ReceiverStream::new(receiver)))
}
RequestBody::Multipart(options) => request.multipart(build_multipart(options)?),
};
}
Expand Down
4 changes: 4 additions & 0 deletions rust/src/transport/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,10 @@ pub enum ConnectionGroup {
#[derive(Debug)]
pub enum RequestBody {
Bytes(Vec<u8>),
Stream {
receiver: UploadReceiver,
length: Option<u64>,
},
Multipart(MultipartBodyOptions),
}

Expand Down
Loading