Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
88 changes: 56 additions & 32 deletions src/asynchronous/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,22 +12,43 @@ use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::{Arc, Mutex};

use async_trait::async_trait;
use tokio::{self, sync::mpsc, task};
use tokio::{
self,
sync::mpsc,
time::{timeout_at, Instant},
};

use crate::error::{get_rpc_status, Error, Result};
use crate::proto::{
Code, Codec, GenMessage, Message, MessageHeader, Request, Response, FLAG_NO_DATA,
FLAG_REMOTE_CLOSED, FLAG_REMOTE_OPEN, MESSAGE_TYPE_DATA, MESSAGE_TYPE_RESPONSE,
};
use crate::r#async::connection::*;
use crate::r#async::shutdown;
use crate::r#async::stream::{
Kind, MessageReceiver, MessageSender, ResultReceiver, ResultSender, StreamInner,
};

use super::stream::SendingMessage;
use super::stream::{MessageControl, SendingMessage};
use super::transport::Socket;

struct StreamRegistrationGuard<'a> {
stream_id: u32,
streams: &'a Mutex<HashMap<u32, ResultSender>>,
}

impl Drop for StreamRegistrationGuard<'_> {
fn drop(&mut self) {
match self.streams.lock() {
Ok(mut streams) => {
streams.remove(&self.stream_id);
}
Err(e) => {
error!("Failed to clean up stream {}: {}", self.stream_id, e);
}
}
}
}

/// A ttrpc Client (async).
#[derive(Clone)]
pub struct Client {
Expand Down Expand Up @@ -76,34 +97,46 @@ impl Client {
/// Requsts a unary request and returns with response.
pub async fn request(&self, req: Request) -> Result<Response> {
let timeout_nano = req.timeout_nano;
let deadline = if timeout_nano == 0 {
None
} else {
Some(Instant::now() + std::time::Duration::from_nanos(timeout_nano as u64))
};
let stream_id = self.next_stream_id.fetch_add(2, Ordering::Relaxed);

let msg: GenMessage = Message::new_request(stream_id, req)?
.try_into()
.map_err(|e: protobuf::Error| Error::Others(e.to_string()))?;

let (tx, mut rx): (ResultSender, ResultReceiver) = mpsc::channel(100);
let control = MessageControl::new(deadline, tx.clone());

self.streams
.lock()
.map_err(|_| Error::Others("Failed to acquire lock on streams".to_string()))?
.insert(stream_id, tx);
let registration = StreamRegistrationGuard {
stream_id,
streams: self.streams.as_ref(),
};

self.req_tx
.send(SendingMessage::new(msg))
.await
.map_err(|_| Error::LocalClosed)?;

let result = if timeout_nano == 0 {
rx.recv().await.ok_or(Error::RemoteClosed)?
let sending_msg = SendingMessage::new_with_control(msg, control);
let send_result = if let Some(deadline) = deadline {
timeout_at(deadline, self.req_tx.send(sending_msg))
.await
.map_err(|_| request_timeout_error())?
} else {
tokio::time::timeout(
std::time::Duration::from_nanos(timeout_nano as u64),
rx.recv(),
)
self.req_tx.send(sending_msg).await
};
send_result.map_err(|_| Error::LocalClosed)?;

let result = if let Some(deadline) = deadline {
timeout_at(deadline, rx.recv())
.await
.map_err(|e| Error::Others(format!("Receive packet timeout {e:?}")))?
.map_err(|_| request_timeout_error())?
.ok_or(Error::RemoteClosed)?
} else {
rx.recv().await.ok_or(Error::RemoteClosed)?
};

let msg = result?;
Expand All @@ -116,6 +149,7 @@ impl Client {
return Err(Error::RpcStatus((*status).clone()));
}

drop(registration);
Ok(res)
}

Expand Down Expand Up @@ -179,16 +213,12 @@ impl Builder for ClientBuilder {
type Writer = ClientWriter;

fn build(&mut self) -> (Self::Reader, Self::Writer) {
let (notifier, waiter) = shutdown::new();
(
ClientReader {
shutdown_waiter: waiter,
streams: self.streams.clone(),
},
ClientWriter {
rx: self.rx.take().unwrap(),
shutdown_notifier: notifier,

streams: self.streams.clone(),
},
)
Expand All @@ -197,8 +227,6 @@ impl Builder for ClientBuilder {

struct ClientWriter {
rx: MessageReceiver,
shutdown_notifier: shutdown::Notifier,

streams: Arc<Mutex<HashMap<u32, ResultSender>>>,
}

Expand Down Expand Up @@ -226,9 +254,7 @@ impl WriterDelegate for ClientWriter {
}
}

async fn exit(&self) {
self.shutdown_notifier.shutdown();
}
async fn exit(&self) {}
}

async fn get_resp_tx(
Expand Down Expand Up @@ -285,21 +311,15 @@ async fn get_resp_tx(

struct ClientReader {
streams: Arc<Mutex<HashMap<u32, ResultSender>>>,
shutdown_waiter: shutdown::Waiter,
}

#[async_trait]
impl ReaderDelegate for ClientReader {
async fn wait_shutdown(&self) {
self.shutdown_waiter.wait_shutdown().await
std::future::pending().await
}

async fn disconnect(&self, e: Error, sender: &mut task::JoinHandle<()>) {
// Abort the request sender task to prevent incoming RPC requests
// from being processed.
sender.abort();
let _ = sender.await;

async fn disconnect(&self, e: Error) {
// Take all items out of `req_map`.
let mut map = std::mem::take(&mut *self.streams.lock().unwrap());
// Terminate undone RPC requests with the error.
Expand Down Expand Up @@ -336,3 +356,7 @@ impl ReaderDelegate for ClientReader {
});
}
}

#[cfg(test)]
#[path = "client_tests.rs"]
mod tests;
Loading
Loading