use std::{
collections::HashMap,
future::Future,
io::{Cursor, ErrorKind},
pin::Pin,
result,
sync::{
Arc,
atomic::{AtomicU32, Ordering},
},
task::{Context, Poll},
};
use async_trait::async_trait;
use bytes::{Buf, BytesMut};
use rmpv::{Value, decode};
#[cfg(feature = "serde")]
use rmpv::{decode::read_value, encode::write_value};
#[cfg(feature = "serde")]
use serde::{Serialize, de::DeserializeOwned};
use tokio::{
io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, WriteHalf, split},
sync::{Mutex, mpsc, oneshot, watch},
task::{JoinError, JoinHandle, JoinSet},
};
use tracing::{error, trace, warn};
use crate::{
error::{ProtocolError, Result, RpcError, ServiceError},
message::*,
};
#[derive(Debug)]
enum ClientMessage {
Request {
id: u32,
method: String,
params: Vec<Value>,
response_sender: oneshot::Sender<Result<Value>>,
},
Notification {
method: String,
params: Vec<Value>,
},
}
#[derive(Debug, Clone)]
pub struct RpcSender {
sender: mpsc::Sender<ClientMessage>,
next_id: Arc<AtomicU32>,
}
impl RpcSender {
fn new(channel: mpsc::Sender<ClientMessage>) -> Self {
Self {
sender: channel,
next_id: Arc::new(AtomicU32::new(1)),
}
}
pub async fn start_request(&self, method: &str, params: &[Value]) -> Result<RequestHandle> {
let id = self.next_id.fetch_add(1, Ordering::Relaxed);
let (response_sender, response_receiver) = oneshot::channel();
self.sender
.send(ClientMessage::Request {
id,
method: method.to_string(),
params: params.to_vec(),
response_sender,
})
.await
.map_err(|_| RpcError::Disconnect { source: None })?;
Ok(RequestHandle {
id,
response: response_receiver,
})
}
pub async fn send_request(&self, method: &str, params: &[Value]) -> Result<Value> {
self.start_request(method, params).await?.response().await
}
pub async fn send_notification(&self, method: &str, params: &[Value]) -> Result<()> {
self.sender
.send(ClientMessage::Notification {
method: method.to_string(),
params: params.to_vec(),
})
.await
.map_err(|_| RpcError::Disconnect { source: None })
}
#[cfg(feature = "serde")]
pub async fn call<Req, Resp>(&self, method: &str, req: &Req) -> Result<Resp>
where
Req: Serialize,
Resp: DeserializeOwned,
{
let params = serialize_params(req)?;
let value = self.send_request(method, ¶ms).await?;
deserialize_response(&value)
}
#[cfg(feature = "serde")]
pub async fn notify<Req>(&self, method: &str, req: &Req) -> Result<()>
where
Req: Serialize,
{
let params = serialize_params(req)?;
self.send_notification(method, ¶ms).await
}
}
#[derive(Debug)]
pub struct RequestHandle {
id: u32,
response: oneshot::Receiver<Result<Value>>,
}
impl RequestHandle {
pub fn id(&self) -> u32 {
self.id
}
pub async fn response(self) -> Result<Value> {
self.response
.await
.map_err(|_| RpcError::Disconnect { source: None })?
}
}
struct AbortOnDrop<T> {
handle: JoinHandle<T>,
}
impl<T> AbortOnDrop<T> {
fn new(handle: JoinHandle<T>) -> Self {
Self { handle }
}
fn abort(&self) {
self.handle.abort();
}
}
impl<T> Drop for AbortOnDrop<T> {
fn drop(&mut self) {
self.abort();
}
}
impl<T> Future for AbortOnDrop<T> {
type Output = result::Result<T, JoinError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
Pin::new(&mut this.handle).poll(cx)
}
}
#[cfg(feature = "serde")]
pub fn serialize_value<T>(value: &T) -> Result<Value>
where
T: Serialize,
{
let buf = rmp_serde::to_vec_named(value)?;
Ok(read_value(&mut Cursor::new(buf))?)
}
#[cfg(feature = "serde")]
pub fn serialize_params<Req>(req: &Req) -> Result<Vec<Value>>
where
Req: Serialize,
{
let value = serialize_value(req)?;
match value {
Value::Array(values) => Ok(values),
value => Ok(vec![value]),
}
}
#[cfg(feature = "serde")]
pub fn deserialize_response<Resp>(value: &Value) -> Result<Resp>
where
Resp: DeserializeOwned,
{
let mut buf = Vec::new();
write_value(&mut buf, value)?;
Ok(rmp_serde::from_slice(&buf)?)
}
#[cfg(feature = "serde")]
pub fn deserialize_params<Req>(params: Vec<Value>) -> Result<Req>
where
Req: DeserializeOwned,
{
let value = Value::Array(params);
deserialize_response(&value)
}
#[cfg(feature = "serde")]
pub fn deserialize_param<Req>(params: Vec<Value>) -> Result<Req>
where
Req: DeserializeOwned,
{
let mut values = params.into_iter();
let value = values
.next()
.ok_or_else(|| RpcError::Protocol(ProtocolError::ExpectedSingleParameter))?;
if values.next().is_some() {
return Err(RpcError::Protocol(ProtocolError::ExpectedSingleParameter));
}
deserialize_response(&value)
}
struct ConnectionHandler<S, T: Connection>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
connection: Arc<Mutex<RpcConnection<S>>>,
service: Arc<T>,
rpc_sender: RpcSender,
}
impl<S, T: Connection> ConnectionHandler<S, T>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
fn new(connection: RpcConnection<S>, service: T, rpc_sender: RpcSender) -> Self {
Self {
connection: Arc::new(Mutex::new(connection)),
service: Arc::new(service),
rpc_sender,
}
}
async fn run(&self, client_receiver: mpsc::Receiver<ClientMessage>) -> Result<()> {
let rpc_sender_clone = self.rpc_sender.clone();
let service = Arc::clone(&self.service);
let mut connected_task = AbortOnDrop::new(tokio::spawn(async move {
service.connected(rpc_sender_clone).await
}));
let mut connected_done = false;
let mut receiver = {
let mut conn = self.connection.lock().await;
conn.take_receiver()?
};
let connection_clone = self.connection.clone();
let client_handler = AbortOnDrop::new(tokio::spawn(async move {
handle_client_messages(connection_clone, client_receiver).await
}));
let mut incoming_handlers: JoinSet<()> = JoinSet::new();
loop {
tokio::select! {
message_result = receiver.recv() => {
match message_result {
Some(Ok(Message::Response(response))) => {
let mut connection = self.connection.lock().await;
if let Err(error) = connection.handle_response(response) {
warn!(%error, "error handling response");
}
}
Some(Ok(message)) => {
let connection = self.connection.clone();
let service = Arc::clone(&self.service);
let rpc_sender = self.rpc_sender.clone();
incoming_handlers.spawn(async move {
if let Err(e) = handle_incoming_message(connection, service, rpc_sender, message).await {
error!("Error handling incoming message: {}", e);
}
});
}
Some(Err(e)) => return Err(e),
None => break,
}
}
connected_result = &mut connected_task, if !connected_done => {
connected_done = true;
match connected_result {
Ok(Ok(())) => {}
Ok(Err(e)) => return Err(e),
Err(source) => {
return Err(RpcError::task_failed("connected callback", source));
}
}
}
Some(joined) = incoming_handlers.join_next(), if !incoming_handlers.is_empty() => {
if let Err(e) = joined
&& !e.is_cancelled() {
error!("Error joining incoming message handler: {}", e);
}
}
else => {
break;
}
}
}
connected_task.abort();
client_handler.abort();
incoming_handlers.abort_all();
while let Some(joined) = incoming_handlers.join_next().await {
if let Err(e) = joined
&& !e.is_cancelled()
{
error!("Error joining incoming message handler: {}", e);
}
}
Ok(())
}
}
async fn handle_incoming_message<S, T>(
connection: Arc<Mutex<RpcConnection<S>>>,
service: Arc<T>,
rpc_sender: RpcSender,
message: Message,
) -> Result<()>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
T: Connection,
{
let service = service.as_ref();
match message {
Message::Request(request) => {
let response = response_from_request_result(
request.id,
service
.handle_request(rpc_sender.clone(), &request.method, request.params)
.await,
);
let mut conn = connection.lock().await;
conn.write_message(&Message::Response(response)).await?;
}
Message::Notification(notification) => {
service
.handle_notification(
rpc_sender.clone(),
¬ification.method,
notification.params,
)
.await?;
}
Message::Response(response) => {
let mut conn = connection.lock().await;
if let Err(e) = conn.handle_response(response) {
warn!("error handling response: {}", e);
}
}
}
Ok(())
}
async fn handle_client_messages<S>(
connection: Arc<Mutex<RpcConnection<S>>>,
mut client_receiver: mpsc::Receiver<ClientMessage>,
) where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
while let Some(message) = client_receiver.recv().await {
let mut conn = connection.lock().await;
let result = match message {
ClientMessage::Request {
id,
method,
params,
response_sender,
} => conn.send_request(id, method, params, response_sender).await,
ClientMessage::Notification { method, params } => {
conn.send_notification(method, params).await
}
};
if let Err(e) = result {
error!("Error handling client message: {}", e);
}
}
}
fn response_from_request_result(id: u32, result: Result<Value>) -> Response {
match result {
Ok(value) => Response {
id,
result: Ok(value),
},
Err(error) => Response {
id,
result: Err(response_error_value(error)),
},
}
}
fn response_error_value(error: RpcError) -> Value {
match error {
RpcError::Service(service_error) => {
warn!("Service error: {}", service_error);
service_error.into()
}
other => {
warn!("RPC error: {}", other);
Value::String(format!("Internal error: {}", other).into())
}
}
}
pub struct ConnectionRuntime<S, T: Connection>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
handler: ConnectionHandler<S, T>,
client_receiver: mpsc::Receiver<ClientMessage>,
rpc_sender: RpcSender,
shutdown_tx: watch::Sender<bool>,
}
impl<S, T> ConnectionRuntime<S, T>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
T: Connection,
{
pub fn new(stream: S, service: T) -> Self {
let connection = RpcConnection::new(stream);
let shutdown_tx = connection.shutdown_sender();
let (sender, client_receiver) = mpsc::channel(100);
let rpc_sender = RpcSender::new(sender);
let handler = ConnectionHandler::new(connection, service, rpc_sender.clone());
Self {
handler,
client_receiver,
rpc_sender,
shutdown_tx,
}
}
pub fn sender(&self) -> RpcSender {
self.rpc_sender.clone()
}
pub fn shutdown_sender(&self) -> watch::Sender<bool> {
self.shutdown_tx.clone()
}
pub async fn run(self) -> Result<()> {
let Self {
handler,
client_receiver,
..
} = self;
handler.run(client_receiver).await
}
}
pub trait ConnectionMaker<T>: Send + Sync
where
T: Connection,
{
fn make_connection(&self) -> T;
}
pub struct ConnectionMakerFn<F> {
make_fn: F,
}
impl<F> ConnectionMakerFn<F> {
pub fn new(make_fn: F) -> Self {
Self { make_fn }
}
}
impl<F, T> ConnectionMaker<T> for ConnectionMakerFn<F>
where
F: Fn() -> T + Send + Sync,
T: Connection,
{
fn make_connection(&self) -> T {
(self.make_fn)()
}
}
impl<T> ConnectionMaker<T> for T
where
T: Connection + Default,
{
fn make_connection(&self) -> T {
Self::default()
}
}
#[async_trait]
pub trait Connection: Send + Sync + 'static {
async fn connected(&self, _client: RpcSender) -> Result<()> {
Ok(())
}
async fn handle_request(
&self,
_client: RpcSender,
method: &str,
params: Vec<Value>,
) -> Result<Value> {
tracing::warn!("Unhandled request: method={}, params={:?}", method, params);
Err(RpcError::Service(ServiceError::method_not_found(method)))
}
async fn handle_notification(
&self,
_client: RpcSender,
method: &str,
params: Vec<Value>,
) -> Result<()> {
tracing::warn!(
"Unhandled notification: method={}, params={:?}",
method,
params
);
Ok(())
}
}
impl Connection for () {}
#[derive(Debug)]
struct RpcConnection<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
message_receiver: Option<mpsc::Receiver<Result<Message>>>,
write_half: WriteHalf<S>,
pending_requests: HashMap<u32, oneshot::Sender<Result<Value>>>,
shutdown_tx: watch::Sender<bool>,
read_task: JoinHandle<()>,
}
impl<S> RpcConnection<S>
where
S: AsyncRead + AsyncWrite + Send + Unpin + 'static,
{
fn new(stream: S) -> Self {
let (read_half, write_half) = split(stream);
let (message_sender, message_receiver) = mpsc::channel(1000);
let (shutdown_tx, mut shutdown_rx) = watch::channel(false);
let read_task = tokio::spawn(async move {
let mut read_half = read_half;
let mut buffer = BytesMut::with_capacity(8192);
let mut eof = false;
loop {
if *shutdown_rx.borrow() {
break;
}
match try_decode_message(&buffer) {
Ok(Some((message, consumed))) => {
buffer.advance(consumed);
if message_sender.send(Ok(message)).await.is_err() {
break;
}
}
Ok(None) if eof => {
drop(
message_sender
.send(Err(RpcError::Disconnect { source: None }))
.await,
);
break;
}
Ok(None) => {
let read_result = tokio::select! {
_ = shutdown_rx.changed() => {
continue;
}
read_result = read_half.read_buf(&mut buffer) => {
read_result
}
};
match read_result {
Ok(0) => {
eof = true;
}
Ok(_) => {}
Err(e) => {
drop(message_sender.send(Err(RpcError::from(e))).await);
break;
}
}
}
Err(e) => {
drop(message_sender.send(Err(e)).await);
break;
}
}
}
});
Self {
write_half,
pending_requests: HashMap::new(),
message_receiver: Some(message_receiver),
shutdown_tx,
read_task,
}
}
fn shutdown_sender(&self) -> watch::Sender<bool> {
self.shutdown_tx.clone()
}
fn take_receiver(&mut self) -> Result<mpsc::Receiver<Result<Message>>> {
self.message_receiver
.take()
.ok_or_else(|| RpcError::resource_already_taken("message receiver"))
}
fn handle_response(&mut self, response: Response) -> Result<()> {
if let Some(sender) = self.pending_requests.remove(&response.id) {
drop(sender.send(response.result.map_err(RpcError::from_remote_error_value)));
Ok(())
} else {
Err(RpcError::Protocol(ProtocolError::UnexpectedResponse {
id: response.id,
}))
}
}
async fn write_message(&mut self, message: &Message) -> Result<()> {
trace!("sending message: {:?}", message);
let mut buffer = Vec::new();
message.encode(&mut buffer)?;
self.write_half.write_all(&buffer).await?;
self.write_half.flush().await?;
Ok(())
}
async fn send_request(
&mut self,
id: u32,
method: String,
params: Vec<Value>,
response_sender: oneshot::Sender<Result<Value>>,
) -> Result<()> {
self.pending_requests.insert(id, response_sender);
let request = Request { id, method, params };
self.write_message(&Message::Request(request)).await
}
async fn send_notification(&mut self, method: String, params: Vec<Value>) -> Result<()> {
let notification = Notification { method, params };
self.write_message(&Message::Notification(notification))
.await
}
}
impl<S> Drop for RpcConnection<S>
where
S: AsyncRead + AsyncWrite + Unpin,
{
fn drop(&mut self) {
self.read_task.abort();
}
}
fn try_decode_message(buffer: &[u8]) -> Result<Option<(Message, usize)>> {
let mut cursor = Cursor::new(buffer);
match decode::read_value(&mut cursor) {
Ok(value) => {
let consumed = cursor.position() as usize;
let message = Message::from_value(value)?;
Ok(Some((message, consumed)))
}
Err(decode::Error::InvalidMarkerRead(e) | decode::Error::InvalidDataRead(e))
if e.kind() == ErrorKind::UnexpectedEof =>
{
Ok(None)
}
Err(decode::Error::DepthLimitExceeded) => {
Err(RpcError::Protocol(ProtocolError::DepthLimitExceeded))
}
Err(e) => Err(RpcError::Deserialization(e)),
}
}
#[cfg(test)]
mod tests {
use tokio::io::duplex;
use super::*;
#[tokio::test]
async fn response_is_resolved_before_following_eof() {
let (client_stream, server_stream) = duplex(1024);
let runtime = ConnectionRuntime::new(client_stream, ());
let sender = runtime.sender();
let runtime_task = tokio::spawn(runtime.run());
let server_task = tokio::spawn(async move {
let mut server = RpcConnection::new(server_stream);
let mut receiver = server.take_receiver().expect("server message receiver");
let Some(Ok(Message::Request(request))) = receiver.recv().await else {
panic!("expected client request");
};
server
.write_message(&Message::Response(Response {
id: request.id,
result: Ok(Value::from(42)),
}))
.await
.expect("write response");
});
let response = sender
.send_request("answer", &[])
.await
.expect("response before disconnect");
assert_eq!(response, Value::from(42));
server_task.await.expect("server task");
assert!(matches!(
runtime_task.await.expect("runtime task"),
Err(RpcError::Disconnect { .. })
));
}
#[test]
fn test_response_from_request_result_preserves_service_errors() {
let response = response_from_request_result(
7,
Err(RpcError::Service(ServiceError::method_not_found("missing"))),
);
assert_eq!(response.id, 7);
assert_eq!(
response.result,
Err(Value::Map(vec![
(
Value::String("name".into()),
Value::String("MethodNotFound".into())
),
(
Value::String("value".into()),
Value::String("Method 'missing' not found".into()),
),
])),
);
}
#[test]
fn test_response_from_request_result_wraps_internal_errors() {
let response =
response_from_request_result(11, Err(RpcError::Protocol("bad request".into())));
assert_eq!(response.id, 11);
assert_eq!(
response.result,
Err(Value::String(
"Internal error: Malformed message: bad request".into(),
)),
);
}
}