co-authored by
Max Brunsfeld
parent
bfccb173c4
commit
0deaa3a61d
+1
-1
@@ -3,4 +3,4 @@ mod peer;
|
||||
pub mod proto;
|
||||
pub mod rest;
|
||||
|
||||
pub use peer::{ConnectionId, Peer, TypedEnvelope};
|
||||
pub use peer::*;
|
||||
|
||||
+84
-19
@@ -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
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user