use std::{
cell::RefCell,
marker::PhantomData,
pin::Pin,
rc::Rc,
task::{Context, Poll},
};
use futures_channel::mpsc;
use futures_core::{future::LocalBoxFuture, Future};
use futures_util::{future::Shared, FutureExt, StreamExt};
use gloo_events::EventListener;
use js_sys::{Array, ArrayBuffer, Uint8Array};
use serde::{de::DeserializeOwned, Deserialize, Serialize};
use wasm_bindgen::JsCast;
#[doc(hidden)]
pub use futures_channel;
#[doc(hidden)]
pub use futures_core;
#[doc(hidden)]
pub use futures_util;
#[doc(hidden)]
pub use gloo_events;
#[doc(hidden)]
pub use js_sys;
#[doc(hidden)]
pub use postcard;
#[doc(hidden)]
pub use postcard_schema;
#[doc(hidden)]
pub use serde;
#[doc(hidden)]
pub use wasm_bindgen;
#[doc(hidden)]
pub use web_sys;
pub use web_rpc_macro::service;
pub mod client;
#[doc(hidden)]
pub mod codec;
pub mod describe;
pub mod interface;
pub mod js;
pub mod port;
#[doc(hidden)]
pub mod service;
pub mod wrap;
pub use interface::Interface;
use port::Port;
#[doc(hidden)]
#[derive(Serialize, Deserialize)]
pub enum MessageHeader {
Request(u32),
Abort(u32),
Response(u32),
StreamItem(u32),
StreamEnd(u32),
}
#[doc(hidden)]
pub type Dispatcher = Shared<LocalBoxFuture<'static, ()>>;
fn to_buffer(bytes: &[u8]) -> ArrayBuffer {
Uint8Array::from(bytes).buffer()
}
#[doc(hidden)]
pub fn take_bytes(message: &Array) -> Vec<u8> {
let buffer = message
.shift()
.dyn_into::<ArrayBuffer>()
.expect("web_rpc: a message must start with an ArrayBuffer");
Uint8Array::new(&buffer).to_vec()
}
#[doc(hidden)]
pub fn post_header(port: &Port, header: MessageHeader) {
let header = to_buffer(&postcard::to_allocvec(&header).unwrap());
let message = Array::of1(&header);
port.post_message(&message, &message).unwrap();
}
#[doc(hidden)]
pub fn post_message(
port: &Port,
header: MessageHeader,
payload: &impl Serialize,
post_args: &Array,
transfer_args: &Array,
) {
let header = to_buffer(&postcard::to_allocvec(&header).unwrap());
let payload = to_buffer(&postcard::to_allocvec(payload).unwrap());
post_args.unshift(&payload);
post_args.unshift(&header);
transfer_args.unshift(&payload);
transfer_args.unshift(&header);
port.post_message(post_args, transfer_args).unwrap();
}
pub struct Builder<C, S> {
client: PhantomData<C>,
service: S,
interface: Interface,
}
impl Builder<(), ()> {
pub fn new(interface: Interface) -> Self {
Self {
interface,
client: PhantomData,
service: (),
}
}
}
impl<C> Builder<C, ()> {
pub fn with_service<S: service::Service>(self, implementation: impl Into<S>) -> Builder<C, S> {
Builder {
interface: self.interface,
client: self.client,
service: implementation.into(),
}
}
}
impl<S> Builder<(), S> {
pub fn with_client<C: client::Client>(self) -> Builder<C, S> {
Builder {
interface: self.interface,
client: PhantomData,
service: self.service,
}
}
}
#[must_use = "Server must be polled in order for RPC requests to be executed"]
pub struct Server {
_listener: Rc<EventListener>,
task: LocalBoxFuture<'static, ()>,
}
impl Future for Server {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.task.poll_unpin(cx)
}
}
struct NoClient;
impl client::Client for NoClient {
type Response = ();
}
impl From<client::State<()>> for NoClient {
fn from(_: client::State<()>) -> Self {
NoClient
}
}
struct NoService;
impl service::Service for NoService {
type Response = ();
async fn execute(
&self,
_: u32,
_: futures_channel::oneshot::Receiver<()>,
_: Vec<u8>,
_: Array,
_: mpsc::UnboundedSender<service::StreamMessage<()>>,
) -> (u32, service::ExecuteResult<()>) {
unreachable!("web_rpc: a request reached an interface with no service")
}
}
fn assemble<C, S>(interface: Interface, service: S) -> (C, Server)
where
C: client::Client + From<client::State<C::Response>> + 'static,
C::Response: DeserializeOwned,
S: service::Service + 'static,
S::Response: Serialize,
{
let Interface {
port,
listener,
mut messages_rx,
} = interface;
let callbacks: Rc<RefCell<client::CallbackMap<C::Response>>> = Default::default();
let stream_callbacks: Rc<RefCell<client::StreamCallbackMap<C::Response>>> = Default::default();
let (requests_tx, requests_rx) = mpsc::unbounded();
let (aborts_tx, aborts_rx) = mpsc::unbounded();
let dispatcher: Dispatcher = {
let callbacks = callbacks.clone();
let stream_callbacks = stream_callbacks.clone();
async move {
while let Some(message) = messages_rx.next().await {
let header: MessageHeader = postcard::from_bytes(&take_bytes(&message)).unwrap();
match header {
MessageHeader::Request(sequence) => {
let payload = take_bytes(&message);
requests_tx
.unbounded_send((sequence, payload, message))
.expect("web_rpc: a request arrived but the server has been dropped");
}
MessageHeader::Abort(sequence) => {
let _ = aborts_tx.unbounded_send(sequence);
}
MessageHeader::Response(sequence) => {
let response = postcard::from_bytes(&take_bytes(&message)).unwrap();
if let Some(callback) = callbacks.borrow_mut().remove(&sequence) {
let _ = callback.send((response, message));
}
}
MessageHeader::StreamItem(sequence) => {
let item = postcard::from_bytes(&take_bytes(&message)).unwrap();
if let Some(items) = stream_callbacks.borrow().get(&sequence) {
let _ = items.unbounded_send((item, message));
}
}
MessageHeader::StreamEnd(sequence) => {
stream_callbacks.borrow_mut().remove(&sequence);
}
}
}
}
.boxed_local()
.shared()
};
let listener = Rc::new(listener);
let client = C::from(client::State {
callbacks,
stream_callbacks,
port: port.clone(),
listener: listener.clone(),
dispatcher: dispatcher.clone(),
sequence: Default::default(),
});
let server = Server {
_listener: listener,
task: service::task::<S>(service, port, dispatcher, requests_rx, aborts_rx).boxed_local(),
};
(client, server)
}
impl<C> Builder<C, ()>
where
C: client::Client + From<client::State<C::Response>> + 'static,
C::Response: DeserializeOwned,
{
pub fn build(self) -> C {
assemble::<C, NoService>(self.interface, NoService).0
}
}
impl<S> Builder<(), S>
where
S: service::Service + 'static,
S::Response: Serialize,
{
pub fn build(self) -> Server {
assemble::<NoClient, S>(self.interface, self.service).1
}
}
impl<C, S> Builder<C, S>
where
C: client::Client + From<client::State<C::Response>> + 'static,
C::Response: DeserializeOwned,
S: service::Service + 'static,
S::Response: Serialize,
{
pub fn build(self) -> (C, Server) {
assemble::<C, S>(self.interface, self.service)
}
}