#![allow(missing_debug_implementations)]
use bytes::Bytes;
use runtime::ConditionallySend;
use wasmrs_frames::{ErrorCode, FrameFlags, RSocketFlags};
use wasmrs_runtime::{self as runtime, unbounded_channel, Entry, SafeMap, UnboundedReceiver, UnboundedSender};
use wasmrs_rx::*;
use crate::{BoxFlux, BoxMono, Frame, PayloadError, RSocket};
mod buffer;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
use futures::stream::{AbortHandle, Abortable};
use futures::{pin_mut, FutureExt, Stream, StreamExt, TryFutureExt};
mod responder;
pub use self::buffer::BufferState;
use self::responder::Responder;
use crate::{Error, RawPayload};
pub enum Handler {
ReqRR(tokio::sync::oneshot::Sender<Result<RawPayload, PayloadError>>),
ReqRS(FluxChannel<RawPayload, PayloadError>),
ReqRC(FluxChannel<RawPayload, PayloadError>),
}
#[derive(Clone, Copy, Debug)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum SocketSide {
Guest,
Host,
}
impl std::fmt::Display for SocketSide {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(match self {
SocketSide::Guest => "guest",
SocketSide::Host => "host",
})
}
}
#[derive()]
#[must_use]
pub struct WasmSocket<T> {
side: SocketSide,
pub(super) handlers: Arc<SafeMap<u32, Handler>>,
abort_handles: Arc<SafeMap<u32, AbortHandle>>,
channels: Arc<SafeMap<u32, UnboundedSender<u32>>>,
pub(super) stream_index: AtomicU32,
tx: UnboundedSender<Frame>,
rx: Option<UnboundedReceiver<Frame>>,
responder: Responder<T>,
n: Arc<AtomicU32>,
}
impl<T: RSocket> std::fmt::Debug for WasmSocket<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ModuleState")
.field("# pending streams", &self.handlers.len())
.field("stream_index", &self.stream_index)
.finish()
}
}
impl<T: RSocket> WasmSocket<T> {
pub fn new(rsocket: T, side: SocketSide) -> WasmSocket<T> {
let first_stream_id = match side {
SocketSide::Guest => 1,
SocketSide::Host => 2,
};
let (snd_tx, snd_rx) = unbounded_channel::<Frame>();
let streams = Arc::new(Default::default());
let abort_handles = Arc::new(Default::default());
let channels = Arc::new(Default::default());
WasmSocket {
side,
stream_index: AtomicU32::new(first_stream_id),
tx: snd_tx,
rx: Some(snd_rx),
handlers: streams,
abort_handles,
channels,
n: Arc::new(AtomicU32::new(Frame::REQUEST_MAX)),
responder: Responder::new(rsocket),
}
}
pub fn set_n(&self, n: u32) {
self.n.store(n, Ordering::SeqCst);
}
pub fn get_n(&self) -> Arc<AtomicU32> {
self.n.clone()
}
pub fn take_rx(&mut self) -> Result<UnboundedReceiver<Frame>, Error> {
self.rx.take().ok_or(crate::Error::ReceiverAlreadyGone)
}
pub(crate) fn next_stream_id(&self) -> u32 {
self.stream_index.fetch_add(2, Ordering::SeqCst)
}
pub fn register_handler(&self, stream_id: u32, handler: Handler) {
self.handlers.insert(stream_id, handler);
}
pub fn register_channel(&self, stream_id: u32) -> UnboundedReceiver<u32> {
let (tx, rx) = unbounded_channel();
self.channels.insert(stream_id, tx);
rx
}
pub fn send(&self, frame: Frame) {
send(&self.tx, self.side, frame);
}
pub fn process_once(&self, frame: Frame) -> Result<(), Error> {
#[cfg(feature = "record-frames")]
crate::record::write_incoming_record(self.side, frame.clone());
let stream_id = frame.stream_id();
trace!(stream_id, side = %self.side, kind = %frame.frame_type(), "process_once");
let flag = frame.get_flag();
match frame {
Frame::RequestFnF(f) => {
let input: RawPayload = f.into();
self.on_request_fnf(stream_id, input);
}
Frame::RequestResponse(f) => {
let input: RawPayload = f.into();
self.on_request_response(stream_id, input);
}
Frame::RequestStream(f) => {
let input: RawPayload = f.into();
self.on_request_stream(stream_id, input);
}
Frame::RequestChannel(f) => {
let input: RawPayload = f.into();
self.on_request_channel(stream_id, input);
}
Frame::PayloadFrame(f) => {
let input: RawPayload = f.into();
self.on_payload(stream_id, flag, input);
}
Frame::Cancel(_) => {
self.on_cancel(stream_id, flag);
}
Frame::ErrorFrame(f) => {
self.on_error(
stream_id,
flag,
f.code,
if f.data.is_empty() {
"Error frame with no data".to_owned()
} else {
f.data
},
f.metadata,
);
}
Frame::RequestN(f) => {
self.on_request_n(stream_id, f.n);
}
}
Ok(())
}
fn on_request_response(&self, sid: u32, input: RawPayload) {
trace!(
sid,
side = %self.side,
"on_request_response"
);
let responder = self.responder.clone();
let tx = self.tx.clone();
let result = responder.request_response(input);
let side = self.side;
runtime::spawn("on_request_response", async move {
match result.await {
Ok(res) => {
send_payload(&tx, sid, side, res, Frame::FLAG_NEXT | Frame::FLAG_COMPLETE);
}
Err(e) => send_error(&tx, sid, side, e),
};
});
}
fn on_request_stream(&self, sid: u32, input: RawPayload) {
trace!(sid, side = %self.side, "on_request_stream");
let responder = self.responder.clone();
let tx = self.tx.clone();
let abort_handles = self.abort_handles.clone();
let side = self.side;
runtime::spawn("on_request_stream", async move {
let (abort_handle, abort_registration) = AbortHandle::new_pair();
abort_handles.insert(sid, abort_handle);
let mut payloads = Abortable::new(responder.request_stream(input), abort_registration);
while let Some(next) = payloads.next().await {
match next {
Ok(it) => send_payload(&tx, sid, side, it, Frame::FLAG_NEXT),
Err(e) => send_error(&tx, sid, side, e),
};
}
abort_handles.remove(&sid);
send_complete(&tx, sid, side, Frame::FLAG_COMPLETE);
});
}
fn on_request_channel(&self, sid: u32, first: RawPayload) {
trace!(sid, side = %self.side, "on_request_channel");
let responder = self.responder.clone();
let tx = self.tx.clone();
let (handler_tx, handler_rx) = FluxChannel::new_parts();
handler_tx.send(first).unwrap();
self.register_handler(sid, Handler::ReqRC(handler_tx));
let abort_handles = self.abort_handles.clone();
let side = self.side;
let n = self.get_n();
runtime::spawn("on_request_channel", async move {
let outputs = responder.request_channel(Box::pin(handler_rx));
let (abort_handle, abort_registration) = AbortHandle::new_pair();
abort_handles.insert(sid, abort_handle);
let mut outputs = Abortable::new(outputs, abort_registration);
let request_n = Frame::new_request_n(sid, n.load(Ordering::SeqCst), 0);
send(&tx, side, request_n);
while let Some(next) = outputs.next().await {
let sending = match next {
Ok(payload) => Frame::new_payload(sid, payload, Frame::FLAG_NEXT),
Err(e) => Frame::new_error(sid, e),
};
send(&tx, side, sending);
}
abort_handles.remove(&sid);
let complete = Frame::new_payload(sid, RawPayload::empty(), Frame::FLAG_COMPLETE);
send(&tx, side, complete);
});
}
fn on_request_fnf(&self, sid: u32, input: RawPayload) {
trace!(sid, side = %self.side, "on_request_fnf");
let responder = self.responder.clone();
let tx = self.tx.clone();
let result = responder.fire_and_forget(input);
let side = self.side;
runtime::spawn("on_request_fnf", async move {
if let Err(e) = result.await {
send_error(&tx, sid, side, e);
}
});
}
fn on_request_n(&self, sid: u32, n: u32) {
trace!(sid, side = %self.side, "on_request_n");
let tx = self.tx.clone();
if n == 0 {
send_error(&tx, sid, self.side, app_err("Invalid RequestN (n=0)"));
return;
}
#[allow(clippy::option_if_let_else)]
match self.channels.cloned(&sid) {
Some(reqn_tx) => {
if reqn_tx.send(n).is_err() {
send_error(&tx, sid, self.side, app_err("RequestN channel closed"));
};
}
None => {
send_error(&tx, sid, self.side, app_err("RequestN called for missing Stream ID"));
}
}
}
fn on_payload(&self, sid: u32, flag: FrameFlags, input: RawPayload) {
trace!(sid, side = %self.side, "on_payload");
let tx = self.tx.clone();
match self.handlers.entry(sid) {
Entry::Occupied(o) => match o.get() {
Handler::ReqRR(_) => match o.remove() {
Handler::ReqRR(sender) => {
if flag.flag_next() && sender.send(Ok(input)).is_err() {
error!(sid, side = %self.side, "error sending payload for REQUEST_RESPONSE, channel already closed");
}
}
_ => unreachable!(),
},
Handler::ReqRS(sender) => {
if flag.flag_next() {
if sender.is_closed() {
warn!(sid, side = %self.side, "request stream already closed");
send_cancel(&tx, sid, self.side);
} else if let Err(_e) = sender.send(input) {
error!(sid, side = %self.side, "error sending payload for REQUEST_STREAM, channel already closed");
send_cancel(&tx, sid, self.side);
}
}
if flag.flag_complete() {
trace!(sid, "removing stream");
o.remove();
}
}
Handler::ReqRC(sender) => {
if flag.flag_next() {
if sender.is_closed() {
warn!(sid, side = %self.side, "request channel already closed");
send_cancel(&tx, sid, self.side);
} else if (sender.send(input)).is_err() {
error!(sid, side = %self.side, "error sending payload for REQUEST_CHANNEL, channel already closed");
send_cancel(&tx, sid, self.side);
}
}
if flag.flag_complete() {
trace!(sid, "removing channel");
o.remove();
}
}
},
Entry::Vacant(_) => {
warn!(sid, side = %self.side, "error sending payload, handler missing");
}
}
}
fn on_cancel(&self, sid: u32, _flag: FrameFlags) {
trace!(sid, side = %self.side, "on_cancel");
if let Some(handler) = self.handlers.remove(&sid) {
let e = PayloadError::new(ErrorCode::Canceled.into(), "Request cancelled", None);
match handler {
Handler::ReqRR(sender) => {
sender.send(Err(e)).unwrap();
}
Handler::ReqRS(_) => {
}
Handler::ReqRC(_) => {
}
}
}
}
fn on_error(&self, sid: u32, flag: FrameFlags, code: u32, message: String, metadata: Option<Bytes>) {
trace!(sid, code, message, ?metadata, side = %self.side, "on_error");
let tx = self.tx.clone();
match self.handlers.entry(sid) {
Entry::Occupied(o) => {
let e = PayloadError::new(code, message, metadata);
match o.get() {
Handler::ReqRR(_) => match o.remove() {
Handler::ReqRR(sender) => {
let _ = sender.send(Err(e));
}
_ => unreachable!(),
},
Handler::ReqRS(sender) => {
if sender.is_closed() {
send_cancel(&tx, sid, self.side);
} else if let Err(_e) = sender.error(e) {
send_cancel(&tx, sid, self.side);
}
if flag.flag_complete() {
o.remove();
}
}
Handler::ReqRC(sender) => {
if sender.is_closed() {
send_cancel(&tx, sid, self.side);
} else if (sender.error(e)).is_err() {
send_cancel(&tx, sid, self.side);
}
if flag.flag_complete() {
o.remove();
}
}
}
}
Entry::Vacant(_) => {}
}
}
}
impl<T: RSocket> RSocket for WasmSocket<T> {
fn fire_and_forget(&self, payload: RawPayload) -> BoxMono<(), PayloadError> {
let sid = self.next_stream_id();
trace!(sid, side = %self.side, "request_response");
let frame = Frame::new_request_fnf(sid, payload, 0);
send(&self.tx, self.side, frame);
futures::future::ready(Ok(())).boxed()
}
fn request_response(&self, payload: RawPayload) -> BoxMono<RawPayload, PayloadError> {
let sid = self.next_stream_id();
trace!(sid, side = %self.side, "request_response");
let (tx, rx) = tokio::sync::oneshot::channel();
self.register_handler(sid, Handler::ReqRR(tx));
let frame = Frame::new_request_response(sid, payload, 0);
send(&self.tx, self.side, frame);
let fut = rx.map_err(|_e| PayloadError::application_error("Request-response channel failed", None));
async move { fut.await? }.boxed()
}
fn request_stream(&self, payload: RawPayload) -> BoxFlux<RawPayload, PayloadError> {
let sid = self.next_stream_id();
trace!(sid, side = %self.side, "request_stream");
let (flux, output) = FluxChannel::new_parts();
self.register_handler(sid, Handler::ReqRS(flux));
let frame = Frame::new_request_stream(sid, payload, 0);
send(&self.tx, self.side, frame);
output.boxed()
}
fn request_channel<S: Stream<Item = Result<RawPayload, PayloadError>> + ConditionallySend + 'static>(
&self,
stream: S,
) -> BoxFlux<RawPayload, PayloadError> {
let sid = self.next_stream_id();
trace!(sid, side = %self.side, "request_channel");
let (tx, rx) = FluxChannel::new_parts();
self.register_handler(sid, Handler::ReqRC(tx));
let mut reqn_rx = self.register_channel(sid);
let tx = self.tx.clone();
let channels = self.channels.clone();
let side = self.side;
runtime::spawn("request_channel", async move {
let mut first = true;
let mut n = 1;
pin_mut!(stream);
while let Some(next) = stream.next().await {
n -= 1;
match next {
Ok(payload) => {
if first {
first = false;
send_channel(&tx, sid, side, payload, Frame::FLAG_NEXT);
} else {
send_payload(&tx, sid, side, payload, Frame::FLAG_NEXT);
}
}
Err(e) => {
send_error(&tx, sid, side, e);
}
}
if n == 0 {
tracing::trace!(%sid,side = %side,"waiting for RequestN");
if let Some(new_n) = reqn_rx.recv().await {
n = new_n;
} else {
break;
}
}
}
channels.remove(&sid);
send_complete(&tx, sid, side, Frame::FLAG_COMPLETE);
trace!(sid, side = %side, "request_channel complete");
});
rx.boxed()
}
}
fn send(tx: &UnboundedSender<Frame>, _side: SocketSide, frame: Frame) {
trace!("sending frame to socket writer: {:?}", frame);
#[cfg(feature = "record-frames")]
crate::record::write_outgoing_record(_side, frame.clone());
if let Err(e) = tx.send(frame) {
warn!(error = %e,side = %_side, "error sending frame to socket writer");
};
}
fn send_payload(tx: &UnboundedSender<Frame>, sid: u32, side: SocketSide, payload: RawPayload, flag: FrameFlags) {
send(tx, side, Frame::new_payload(sid, payload, flag));
}
fn send_channel(tx: &UnboundedSender<Frame>, sid: u32, side: SocketSide, payload: RawPayload, flag: FrameFlags) {
send(
tx,
side,
Frame::new_request_channel(sid, payload, flag, Frame::REQUEST_MAX),
);
}
fn send_cancel(tx: &UnboundedSender<Frame>, sid: u32, side: SocketSide) {
send(tx, side, Frame::new_cancel(sid));
}
fn send_complete(tx: &UnboundedSender<Frame>, sid: u32, side: SocketSide, flag: FrameFlags) {
send(tx, side, Frame::new_payload(sid, RawPayload::empty(), flag));
}
fn send_error(tx: &UnboundedSender<Frame>, sid: u32, side: SocketSide, e: PayloadError) {
let error = Frame::new_error(sid, e);
send(tx, side, error);
}
fn app_err(msg: impl AsRef<str>) -> PayloadError {
PayloadError::application_error(msg.as_ref(), None)
}
#[cfg(test)]
mod test {
use anyhow::Result;
use bytes::Bytes;
use super::*;
#[derive(Clone)]
struct EchoRSocket;
impl RSocket for EchoRSocket {
fn fire_and_forget(&self, _payload: RawPayload) -> BoxMono<(), PayloadError> {
futures::future::ready(Ok(())).boxed()
}
fn request_response(&self, payload: RawPayload) -> BoxMono<RawPayload, PayloadError> {
info!("{:?}", payload);
futures::future::ready(Ok(payload)).boxed()
}
fn request_stream(&self, payload: RawPayload) -> BoxFlux<RawPayload, PayloadError> {
info!("{:?}", payload);
let (tx, rx) = FluxChannel::new_parts();
tx.send(payload.clone()).unwrap();
tx.send(payload).unwrap();
tx.complete();
Box::pin(rx)
}
fn request_channel<T: Stream<Item = std::result::Result<RawPayload, PayloadError>> + Send + 'static>(
&self,
stream: T,
) -> BoxFlux<RawPayload, PayloadError> {
let (tx, rx) = FluxChannel::new_parts();
runtime::spawn("request_channel", async move {
pin_mut!(stream);
while let Some(next) = stream.next().await {
tx.send_result(next).unwrap();
}
tx.complete();
});
Box::pin(rx)
}
}
fn make_echo() -> (Arc<WasmSocket<EchoRSocket>>, Arc<WasmSocket<EchoRSocket>>) {
let mut guest = WasmSocket::new(EchoRSocket {}, SocketSide::Guest);
let mut guest_frame_rx = guest.take_rx().unwrap();
let mut host = WasmSocket::new(EchoRSocket {}, SocketSide::Host);
let mut host_frame_rx = host.take_rx().unwrap();
let guest = Arc::new(guest);
let inner_guest = guest.clone();
let host = Arc::new(host);
let inner_host = host.clone();
runtime::spawn("guest->host", async move {
while let Some(frame) = guest_frame_rx.recv().await {
println!("GUEST >>> HOST: {:?}", frame);
inner_host.process_once(frame).unwrap();
}
});
runtime::spawn("host->guest", async move {
while let Some(frame) = host_frame_rx.recv().await {
println!("HOST >>> GUEST: {:?}", frame);
inner_guest.process_once(frame).unwrap();
}
});
(guest, host)
}
#[test_log::test(tokio::test)]
async fn test_fnf() -> Result<()> {
let (guest, _host) = make_echo();
let output = guest
.fire_and_forget(RawPayload::new(Bytes::from_static(b""), Bytes::from_static(b"FNF")))
.await;
assert!(output.is_ok());
Ok(())
}
#[test_log::test(tokio::test)]
async fn test_reqres() -> Result<()> {
let (guest, _host) = make_echo();
let output = guest.request_response(RawPayload::new(Bytes::from_static(b""), Bytes::from_static(b"REQRES")));
let once = output.await.unwrap();
assert_eq!(once.data, Some(Bytes::from_static(b"REQRES")));
Ok(())
}
#[test_log::test(tokio::test)]
async fn test_reqstream() -> Result<()> {
let (guest, _host) = make_echo();
let mut output = guest.request_stream(RawPayload::new(Bytes::from_static(b""), Bytes::from_static(b"REQ_STR")));
let once = output.next().await.unwrap().unwrap();
assert_eq!(once.data, Some(Bytes::from_static(b"REQ_STR")));
let once = output.next().await.unwrap().unwrap();
assert_eq!(once.data, Some(Bytes::from_static(b"REQ_STR")));
Ok(())
}
#[test_log::test(tokio::test)]
async fn test_reqchannel() -> Result<()> {
let (guest, _host) = make_echo();
let (tx, rx) = FluxChannel::new_parts();
let mut output = guest.request_channel(Box::pin(rx));
tx.send(RawPayload::new(
Bytes::from_static(b""),
Bytes::from_static(b"REQCHANNEL1"),
))
.unwrap();
tx.send(RawPayload::new(
Bytes::from_static(b""),
Bytes::from_static(b"REQCHANNEL2"),
))
.unwrap();
tx.complete();
let once = output.next().await.unwrap().unwrap();
assert_eq!(once.data, Some(Bytes::from_static(b"REQCHANNEL1")));
let once = output.next().await.unwrap().unwrap();
assert_eq!(once.data, Some(Bytes::from_static(b"REQCHANNEL2")));
Ok(())
}
}