Co-Authored-By: Max Brunsfeld <maxbrunsfeld@gmail.com>
This commit is contained in:
Nathan Sobo
2021-06-18 17:26:12 -06:00
co-authored by Max Brunsfeld
parent bfccb173c4
commit 0deaa3a61d
7 changed files with 235 additions and 62 deletions
+1 -1
View File
@@ -3,4 +3,4 @@ mod peer;
pub mod proto;
pub mod rest;
pub use peer::{ConnectionId, Peer, TypedEnvelope};
pub use peer::*;
+84 -19
View File
@@ -14,6 +14,7 @@ use std::{
collections::{HashMap, HashSet},
fmt,
future::Future,
marker::PhantomData,
pin::Pin,
sync::{
atomic::{self, AtomicU32},
@@ -27,6 +28,9 @@ type BoxedReader = Pin<Box<dyn AsyncRead + 'static + Send>>;
#[derive(Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub struct ConnectionId(u32);
#[derive(Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub struct PeerId(u32);
struct Connection {
writer: Mutex<MessageStream<BoxedWriter>>,
reader: Mutex<MessageStream<BoxedReader>>,
@@ -38,12 +42,30 @@ type MessageHandler = Box<
dyn Send + Sync + Fn(&mut Option<proto::Envelope>, ConnectionId) -> Option<BoxFuture<bool>>,
>;
#[derive(Clone, Copy)]
pub struct Receipt<T> {
sender_id: ConnectionId,
message_id: u32,
payload_type: PhantomData<T>,
}
pub struct TypedEnvelope<T> {
pub id: u32,
pub connection_id: ConnectionId,
pub sender_id: ConnectionId,
pub original_sender_id: Option<PeerId>,
pub message_id: u32,
pub payload: T,
}
impl<T: RequestMessage> TypedEnvelope<T> {
pub fn receipt(&self) -> Receipt<T> {
Receipt {
sender_id: self.sender_id,
message_id: self.message_id,
payload_type: PhantomData,
}
}
}
pub struct Peer {
connections: RwLock<HashMap<ConnectionId, Arc<Connection>>>,
connection_close_barriers: RwLock<HashMap<ConnectionId, barrier::Sender>>,
@@ -81,8 +103,9 @@ impl Peer {
Some(
async move {
tx.send(TypedEnvelope {
id: envelope.id,
connection_id,
sender_id: connection_id,
original_sender_id: envelope.original_sender_id.map(PeerId),
message_id: envelope.id,
payload: T::from_envelope(envelope).unwrap(),
})
.await
@@ -200,25 +223,45 @@ impl Peer {
) -> Result<TypedEnvelope<M>> {
let connection = self.connection(connection_id).await?;
let envelope = connection.reader.lock().await.read_message().await?;
let id = envelope.id;
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 {
id,
connection_id,
sender_id: connection_id,
original_sender_id: original_sender_id.map(PeerId),
message_id,
payload,
})
}
pub fn request<T: RequestMessage>(
self: &Arc<Self>,
connection_id: ConnectionId,
req: T,
receiver_id: ConnectionId,
request: T,
) -> impl Future<Output = Result<T::Response>> {
self.request_internal(None, receiver_id, request)
}
pub fn forward_request<T: RequestMessage>(
self: &Arc<Self>,
sender_id: ConnectionId,
receiver_id: ConnectionId,
request: T,
) -> impl Future<Output = Result<T::Response>> {
self.request_internal(Some(sender_id), receiver_id, request)
}
pub fn request_internal<T: RequestMessage>(
self: &Arc<Self>,
original_sender_id: Option<ConnectionId>,
receiver_id: ConnectionId,
request: T,
) -> impl Future<Output = Result<T::Response>> {
let this = self.clone();
let (tx, mut rx) = oneshot::channel();
async move {
let connection = this.connection(connection_id).await?;
let connection = this.connection(receiver_id).await?;
let message_id = connection
.next_message_id
.fetch_add(1, atomic::Ordering::SeqCst);
@@ -231,7 +274,11 @@ impl Peer {
.writer
.lock()
.await
.write_message(&req.into_envelope(message_id, None))
.write_message(&request.into_envelope(
message_id,
None,
original_sender_id.map(|id| id.0),
))
.await?;
let response = rx
.recv()
@@ -257,7 +304,7 @@ impl Peer {
.writer
.lock()
.await
.write_message(&message.into_envelope(message_id, None))
.write_message(&message.into_envelope(message_id, None, None))
.await?;
Ok(())
}
@@ -265,12 +312,12 @@ impl Peer {
pub fn respond<T: RequestMessage>(
self: &Arc<Self>,
request: TypedEnvelope<T>,
receipt: Receipt<T>,
response: T::Response,
) -> impl Future<Output = Result<()>> {
let this = self.clone();
async move {
let connection = this.connection(request.connection_id).await?;
let connection = this.connection(receipt.sender_id).await?;
let message_id = connection
.next_message_id
.fetch_add(1, atomic::Ordering::SeqCst);
@@ -278,7 +325,7 @@ impl Peer {
.writer
.lock()
.await
.write_message(&response.into_envelope(message_id, Some(request.id)))
.write_message(&response.into_envelope(message_id, Some(receipt.message_id), None))
.await?;
Ok(())
}
@@ -301,6 +348,12 @@ impl fmt::Display for ConnectionId {
}
}
impl fmt::Display for PeerId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(f)
}
}
#[cfg(test)]
mod tests {
use super::*;
@@ -396,19 +449,31 @@ mod tests {
async move {
let msg = auth_rx.recv().await.unwrap();
assert_eq!(msg.payload, request1);
server.respond(msg, response1.clone()).await.unwrap();
server
.respond(msg.receipt(), response1.clone())
.await
.unwrap();
let msg = auth_rx.recv().await.unwrap();
assert_eq!(msg.payload, request2.clone());
server.respond(msg, response2.clone()).await.unwrap();
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, response3.clone()).await.unwrap();
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, response4.clone()).await.unwrap();
server
.respond(msg.receipt(), response4.clone())
.await
.unwrap();
server_done_tx.send(()).await.unwrap();
}
+15 -4
View File
@@ -6,7 +6,12 @@ include!(concat!(env!("OUT_DIR"), "/zed.messages.rs"));
pub trait EnvelopedMessage: Sized + Send + 'static {
const NAME: &'static str;
fn into_envelope(self, id: u32, responding_to: Option<u32>) -> Envelope;
fn into_envelope(
self,
id: u32,
responding_to: Option<u32>,
original_sender_id: Option<u32>,
) -> Envelope;
fn matches_envelope(envelope: &Envelope) -> bool;
fn from_envelope(envelope: Envelope) -> Option<Self>;
}
@@ -20,10 +25,16 @@ macro_rules! message {
impl EnvelopedMessage for $name {
const NAME: &'static str = std::stringify!($name);
fn into_envelope(self, id: u32, responding_to: Option<u32>) -> Envelope {
fn into_envelope(
self,
id: u32,
responding_to: Option<u32>,
original_sender_id: Option<u32>,
) -> Envelope {
Envelope {
id,
responding_to,
original_sender_id,
payload: Some(envelope::Payload::$name(self)),
}
}
@@ -132,13 +143,13 @@ mod tests {
user_id: 5,
access_token: "the-access-token".into(),
}
.into_envelope(3, None);
.into_envelope(3, None, None);
let message2 = OpenBuffer {
worktree_id: 1,
path: "path".to_string(),
}
.into_envelope(5, None);
.into_envelope(5, None, None);
let mut message_stream = MessageStream::new(byte_stream);
message_stream.write_message(&message1).await.unwrap();