use crate::proto::{self, EnvelopedMessage, MessageStream, RequestMessage}; use anyhow::{anyhow, Context, Result}; use async_lock::{Mutex, RwLock}; use futures::{future::BoxFuture, AsyncRead, AsyncWrite, FutureExt}; use postage::{ mpsc, prelude::{Sink, Stream}, }; use std::{ any::TypeId, collections::{HashMap, HashSet}, fmt, future::Future, marker::PhantomData, sync::{ atomic::{self, AtomicU32}, Arc, }, }; #[derive(Clone, Copy, PartialEq, Eq, Hash, Debug)] pub struct ConnectionId(pub u32); #[derive(Clone, Copy, PartialEq, Eq, Hash, Debug)] pub struct PeerId(pub u32); type MessageHandler = Box< dyn Send + Sync + Fn(&mut Option, ConnectionId) -> Option>, >; pub struct Receipt { sender_id: ConnectionId, message_id: u32, payload_type: PhantomData, } pub struct TypedEnvelope { pub sender_id: ConnectionId, original_sender_id: Option, pub message_id: u32, pub payload: T, } impl TypedEnvelope { pub fn original_sender_id(&self) -> Result { self.original_sender_id .ok_or_else(|| anyhow!("missing original_sender_id")) } } impl TypedEnvelope { pub fn receipt(&self) -> Receipt { Receipt { sender_id: self.sender_id, message_id: self.message_id, payload_type: PhantomData, } } } pub struct Peer { connections: RwLock>, message_handlers: RwLock>, handler_types: Mutex>, next_connection_id: AtomicU32, } #[derive(Clone)] struct Connection { outgoing_tx: mpsc::Sender, next_message_id: Arc, response_channels: ResponseChannels, } pub struct ConnectionHandler { peer: Arc, connection_id: ConnectionId, response_channels: ResponseChannels, outgoing_rx: mpsc::Receiver, reader: MessageStream, writer: MessageStream, } type ResponseChannels = Arc>>>; impl Peer { pub fn new() -> Arc { Arc::new(Self { connections: Default::default(), message_handlers: Default::default(), handler_types: Default::default(), next_connection_id: Default::default(), }) } pub async fn add_message_handler( &self, ) -> mpsc::Receiver> { if !self.handler_types.lock().await.insert(TypeId::of::()) { panic!("duplicate handler type"); } let (tx, rx) = mpsc::channel(256); self.message_handlers .write() .await .push(Box::new(move |envelope, connection_id| { if envelope.as_ref().map_or(false, T::matches_envelope) { let envelope = Option::take(envelope).unwrap(); let mut tx = tx.clone(); Some( async move { tx.send(TypedEnvelope { sender_id: connection_id, original_sender_id: envelope.original_sender_id.map(PeerId), message_id: envelope.id, payload: T::from_envelope(envelope).unwrap(), }) .await .is_err() } .boxed(), ) } else { None } })); rx } pub async fn add_connection( self: &Arc, conn: Conn, ) -> (ConnectionId, ConnectionHandler) where Conn: Clone + AsyncRead + AsyncWrite + Unpin + Send + 'static, { let connection_id = ConnectionId( self.next_connection_id .fetch_add(1, atomic::Ordering::SeqCst), ); let (outgoing_tx, outgoing_rx) = mpsc::channel(64); let connection = Connection { outgoing_tx, next_message_id: Default::default(), response_channels: Default::default(), }; let handler = ConnectionHandler { peer: self.clone(), connection_id, response_channels: connection.response_channels.clone(), outgoing_rx, reader: MessageStream::new(conn.clone()), writer: MessageStream::new(conn), }; self.connections .write() .await .insert(connection_id, connection); (connection_id, handler) } pub async fn disconnect(&self, connection_id: ConnectionId) { self.connections.write().await.remove(&connection_id); } pub async fn reset(&self) { self.connections.write().await.clear(); self.handler_types.lock().await.clear(); self.message_handlers.write().await.clear(); } pub fn request( self: &Arc, receiver_id: ConnectionId, request: T, ) -> impl Future> { self.request_internal(None, receiver_id, request) } pub fn forward_request( self: &Arc, sender_id: ConnectionId, receiver_id: ConnectionId, request: T, ) -> impl Future> { self.request_internal(Some(sender_id), receiver_id, request) } pub fn request_internal( self: &Arc, original_sender_id: Option, receiver_id: ConnectionId, request: T, ) -> impl Future> { let this = self.clone(); let (tx, mut rx) = mpsc::channel(1); async move { let mut connection = this.connection(receiver_id).await?; let message_id = connection .next_message_id .fetch_add(1, atomic::Ordering::SeqCst); connection .response_channels .lock() .await .insert(message_id, tx); connection .outgoing_tx .send(request.into_envelope(message_id, None, original_sender_id.map(|id| id.0))) .await?; let response = rx .recv() .await .ok_or_else(|| anyhow!("connection was closed"))?; T::Response::from_envelope(response) .ok_or_else(|| anyhow!("received response of the wrong type")) } } pub fn send( self: &Arc, receiver_id: ConnectionId, message: T, ) -> impl Future> { let this = self.clone(); async move { let mut connection = this.connection(receiver_id).await?; let message_id = connection .next_message_id .fetch_add(1, atomic::Ordering::SeqCst); connection .outgoing_tx .send(message.into_envelope(message_id, None, None)) .await?; Ok(()) } } pub fn forward_send( self: &Arc, sender_id: ConnectionId, receiver_id: ConnectionId, message: T, ) -> impl Future> { let this = self.clone(); async move { let mut connection = this.connection(receiver_id).await?; let message_id = connection .next_message_id .fetch_add(1, atomic::Ordering::SeqCst); connection .outgoing_tx .send(message.into_envelope(message_id, None, Some(sender_id.0))) .await?; Ok(()) } } pub fn respond( self: &Arc, receipt: Receipt, response: T::Response, ) -> impl Future> { let this = self.clone(); async move { let mut connection = this.connection(receipt.sender_id).await?; let message_id = connection .next_message_id .fetch_add(1, atomic::Ordering::SeqCst); connection .outgoing_tx .send(response.into_envelope(message_id, Some(receipt.message_id), None)) .await?; Ok(()) } } fn connection( self: &Arc, connection_id: ConnectionId, ) -> impl Future> { let this = self.clone(); async move { let connections = this.connections.read().await; let connection = connections .get(&connection_id) .ok_or_else(|| anyhow!("no such connection: {}", connection_id))?; Ok(connection.clone()) } } } impl ConnectionHandler where Conn: Clone + AsyncRead + AsyncWrite + Unpin + Send + 'static, { pub async fn run(mut self) -> Result<()> { loop { let read_message = self.reader.read_message().fuse(); futures::pin_mut!(read_message); loop { futures::select! { incoming = read_message => match incoming { Ok(incoming) => { Self::handle_incoming_message(incoming, &self.peer, self.connection_id, &self.response_channels).await; break; } Err(error) => { self.response_channels.lock().await.clear(); Err(error).context("received invalid RPC message")?; } }, outgoing = self.outgoing_rx.recv().fuse() => match outgoing { Some(outgoing) => { if let Err(result) = self.writer.write_message(&outgoing).await { self.response_channels.lock().await.clear(); Err(result).context("failed to write RPC message")?; } } None => return Ok(()), } } } } } pub async fn receive(&mut self) -> Result> { let envelope = self.reader.read_message().await?; let original_sender_id = envelope.original_sender_id; let message_id = envelope.id; let payload = M::from_envelope(envelope).ok_or_else(|| anyhow!("unexpected message type"))?; Ok(TypedEnvelope { sender_id: self.connection_id, original_sender_id: original_sender_id.map(PeerId), message_id, payload, }) } async fn handle_incoming_message( message: proto::Envelope, peer: &Arc, connection_id: ConnectionId, response_channels: &ResponseChannels, ) { if let Some(responding_to) = message.responding_to { let channel = response_channels.lock().await.remove(&responding_to); if let Some(mut tx) = channel { tx.send(message).await.ok(); } else { log::warn!("received RPC response to unknown request {}", responding_to); } } else { let mut envelope = Some(message); let mut handler_index = None; let mut handler_was_dropped = false; for (i, handler) in peer.message_handlers.read().await.iter().enumerate() { if let Some(future) = handler(&mut envelope, connection_id) { handler_was_dropped = future.await; handler_index = Some(i); break; } } if let Some(handler_index) = handler_index { if handler_was_dropped { drop(peer.message_handlers.write().await.remove(handler_index)); } } else { log::warn!("unhandled message: {:?}", envelope.unwrap().payload); } } } } impl Clone for Receipt { fn clone(&self) -> Self { Self { sender_id: self.sender_id, message_id: self.message_id, payload_type: PhantomData, } } } impl Copy for Receipt {} impl fmt::Display for ConnectionId { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { self.0.fmt(f) } } impl fmt::Display for PeerId { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { self.0.fmt(f) } } #[cfg(test)] mod tests { use super::*; use postage::oneshot; use smol::{ io::AsyncWriteExt, net::unix::{UnixListener, UnixStream}, }; use std::io; use tempdir::TempDir; #[test] fn test_request_response() { smol::block_on(async move { // create socket let socket_dir_path = TempDir::new("test-request-response").unwrap(); let socket_path = socket_dir_path.path().join("test.sock"); let listener = UnixListener::bind(&socket_path).unwrap(); // create 2 clients connected to 1 server let server = Peer::new(); let client1 = Peer::new(); let client2 = Peer::new(); let (client1_conn_id, task1) = client1 .add_connection(UnixStream::connect(&socket_path).await.unwrap()) .await; let (client2_conn_id, task2) = client2 .add_connection(UnixStream::connect(&socket_path).await.unwrap()) .await; let (_, task3) = server .add_connection(listener.accept().await.unwrap().0) .await; let (_, task4) = server .add_connection(listener.accept().await.unwrap().0) .await; smol::spawn(task1.run()).detach(); smol::spawn(task2.run()).detach(); smol::spawn(task3.run()).detach(); smol::spawn(task4.run()).detach(); // define the expected requests and responses let request1 = proto::Auth { user_id: 1, access_token: "token-1".to_string(), }; let response1 = proto::AuthResponse { credentials_valid: true, }; let request2 = proto::Auth { user_id: 2, access_token: "token-2".to_string(), }; let response2 = proto::AuthResponse { credentials_valid: false, }; let request3 = proto::OpenBuffer { worktree_id: 1, path: "path/two".to_string(), }; let response3 = proto::OpenBufferResponse { buffer: Some(proto::Buffer { id: 2, content: "path/two content".to_string(), history: vec![], selections: vec![], }), }; let request4 = proto::OpenBuffer { worktree_id: 2, path: "path/one".to_string(), }; let response4 = proto::OpenBufferResponse { buffer: Some(proto::Buffer { id: 1, content: "path/one content".to_string(), history: vec![], selections: vec![], }), }; // on the server, respond to two requests for each client let mut open_buffer_rx = server.add_message_handler::().await; let mut auth_rx = server.add_message_handler::().await; let (mut server_done_tx, mut server_done_rx) = oneshot::channel::<()>(); smol::spawn({ let request1 = request1.clone(); let request2 = request2.clone(); let request3 = request3.clone(); let request4 = request4.clone(); let response1 = response1.clone(); let response2 = response2.clone(); let response3 = response3.clone(); let response4 = response4.clone(); async move { let msg = auth_rx.recv().await.unwrap(); assert_eq!(msg.payload, request1); server .respond(msg.receipt(), response1.clone()) .await .unwrap(); let msg = auth_rx.recv().await.unwrap(); assert_eq!(msg.payload, request2.clone()); server .respond(msg.receipt(), response2.clone()) .await .unwrap(); let msg = open_buffer_rx.recv().await.unwrap(); assert_eq!(msg.payload, request3.clone()); server .respond(msg.receipt(), response3.clone()) .await .unwrap(); let msg = open_buffer_rx.recv().await.unwrap(); assert_eq!(msg.payload, request4.clone()); server .respond(msg.receipt(), response4.clone()) .await .unwrap(); server_done_tx.send(()).await.unwrap(); } }) .detach(); assert_eq!( client1.request(client1_conn_id, request1).await.unwrap(), response1 ); assert_eq!( client2.request(client2_conn_id, request2).await.unwrap(), response2 ); assert_eq!( client2.request(client2_conn_id, request3).await.unwrap(), response3 ); assert_eq!( client1.request(client1_conn_id, request4).await.unwrap(), response4 ); client1.disconnect(client1_conn_id).await; client2.disconnect(client1_conn_id).await; server_done_rx.recv().await.unwrap(); }); } #[test] fn test_disconnect() { smol::block_on(async move { let socket_dir_path = TempDir::new("drop-client").unwrap(); let socket_path = socket_dir_path.path().join(".sock"); let listener = UnixListener::bind(&socket_path).unwrap(); let client_conn = UnixStream::connect(&socket_path).await.unwrap(); let (mut server_conn, _) = listener.accept().await.unwrap(); let client = Peer::new(); let (connection_id, handler) = client.add_connection(client_conn).await; let (mut incoming_messages_ended_tx, mut incoming_messages_ended_rx) = postage::barrier::channel(); smol::spawn(async move { handler.run().await.ok(); incoming_messages_ended_tx.send(()).await.unwrap(); }) .detach(); client.disconnect(connection_id).await; incoming_messages_ended_rx.recv().await; let err = server_conn.write(&[]).await.unwrap_err(); assert_eq!(err.kind(), io::ErrorKind::BrokenPipe); }); } #[test] fn test_io_error() { smol::block_on(async move { let socket_dir_path = TempDir::new("io-error").unwrap(); let socket_path = socket_dir_path.path().join(".sock"); let _listener = UnixListener::bind(&socket_path).unwrap(); let mut client_conn = UnixStream::connect(&socket_path).await.unwrap(); client_conn.close().await.unwrap(); let client = Peer::new(); let (connection_id, handler) = client.add_connection(client_conn).await; smol::spawn(handler.run()).detach(); let err = client .request( connection_id, proto::Auth { user_id: 42, access_token: "token".to_string(), }, ) .await .unwrap_err(); assert_eq!(err.to_string(), "connection was closed"); }); } }