Skip to content
Merged
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
16 changes: 11 additions & 5 deletions example-prost/async_client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -35,11 +35,17 @@ async fn main() {
)
.await;

assert_eq!(
resp,
Err(ttrpc::Error::Others(
"Receive packet timeout Elapsed(())".into()
))
// Either peer can report the deadline first, depending on scheduling.
assert!(
match &resp {
Err(ttrpc::Error::Others(message)) => message == "Request deadline elapsed",
Err(ttrpc::Error::RpcStatus(status)) => {
status.code() == ttrpc::Code::DEADLINE_EXCEEDED
}
_ => false,
},
"expected a deadline error, got {:?}",
resp
);
println!(
"Green Thread 1 - {} -> {:?} ended: {:?}",
Expand Down
16 changes: 11 additions & 5 deletions example/async-client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -37,11 +37,17 @@ async fn main() {
)
.await;

assert_eq!(
resp,
Err(ttrpc::Error::Others(
"Receive packet timeout Elapsed(())".into()
))
// Either peer can report the deadline first, depending on scheduling.
assert!(
match &resp {
Err(ttrpc::Error::Others(message)) => message == "Request deadline elapsed",
Err(ttrpc::Error::RpcStatus(status)) => {
status.code() == ttrpc::Code::DEADLINE_EXCEEDED
}
_ => false,
},
"expected a deadline error, got {:?}",
resp
);

println!(
Expand Down
123 changes: 74 additions & 49 deletions src/asynchronous/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,11 @@ 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::{
Expand All @@ -21,17 +25,44 @@ use crate::proto::{
MESSAGE_TYPE_RESPONSE,
};
use crate::r#async::connection::*;
use crate::r#async::shutdown;
use crate::r#async::stream::{
ClientResultSender, ClientStreams, MessageReceiver, MessageSender, StreamInner,
};
#[cfg(feature = "security_extension")]
use crate::security_extension::ConnectHook;
use crate::ConnectionContext;

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

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

impl StreamRegistrationGuard<'_> {
fn disarm(mut self) {
self.active = false;
}
}

impl Drop for StreamRegistrationGuard<'_> {
fn drop(&mut self) {
if !self.active {
return;
}
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 cloneable asynchronous ttrpc connection.
///
/// Generated service clients wrap this type. Clones share one connection and can issue concurrent
Expand Down Expand Up @@ -189,6 +220,11 @@ impl Client {
/// timeout expires, the response is malformed, or the server returns a non-OK status.
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 mut msg: GenMessage = Message::new_request(stream_id, req)?
Expand All @@ -200,32 +236,31 @@ impl Client {
check_oversize(msg.payload.len(), false)?;

let (tx, mut rx) = mpsc::unbounded_channel();
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(),
active: true,
};

// ── Injection Point 6/10: unary REQUEST transform_outbound ──
if let Err(e) = self
.conn_ctx
.transform_send(&mut msg, &self.req_tx, false, false)
.await
{
self.streams.lock().unwrap().remove(&stream_id);
return Err(e);
}
self.conn_ctx
.transform_send_with_control(&mut msg, &self.req_tx, false, false, control)
.await?;

let result = if timeout_nano == 0 {
rx.recv().await.ok_or(Error::RemoteClosed)?
} else {
tokio::time::timeout(
std::time::Duration::from_nanos(timeout_nano as u64),
rx.recv(),
)
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)?
};
registration.disarm();

let msg = result?;

Expand Down Expand Up @@ -281,26 +316,28 @@ impl Client {
.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(),
active: true,
};

// ── Injection Point 8/10: stream-init REQUEST transform_outbound ──
if let Err(e) = self
.conn_ctx
self.conn_ctx
.transform_send(&mut msg, &self.req_tx, false, false)
.await
{
self.streams.lock().unwrap().remove(&stream_id);
return Err(e);
}
.await?;

Ok(StreamInner::new_client(
let inner = StreamInner::new_client(
stream_id,
self.req_tx.clone(),
rx,
streaming_client,
streaming_server,
self.streams.clone(),
self.conn_ctx.clone(),
))
);
registration.disarm();
Ok(inner)
}
}

Expand All @@ -316,24 +353,20 @@ 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(),
conn_ctx: self.conn_ctx.clone(),
},
ClientWriter {
rx: self.rx.take().unwrap(),
shutdown_notifier: notifier,
},
)
}
}

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

#[async_trait]
Expand All @@ -342,29 +375,21 @@ impl WriterDelegate for ClientWriter {
self.rx.recv().await
}

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

struct ClientReader {
streams: ClientStreams,
shutdown_waiter: shutdown::Waiter,
conn_ctx: Arc<ConnectionContext>,
}

#[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 every pending RPC with the error. Enqueuing into each
Expand Down Expand Up @@ -451,7 +476,7 @@ impl ClientReader {
}

#[cfg(all(test, feature = "security_extension"))]
mod tests {
mod security_tests {
use super::*;
use crate::security_extension::{ConnectHook, ConnectionData, HookError, HookOutput};

Expand Down Expand Up @@ -578,18 +603,14 @@ mod teardown_tests {
let (tx2, mut rx2) = mpsc::unbounded_channel();
streams.lock().unwrap().insert(2, tx2);

let (_notifier, waiter) = shutdown::new();
let reader = ClientReader {
streams: streams.clone(),
shutdown_waiter: waiter,
conn_ctx: Arc::new(ConnectionContext::default()),
};

let mut dummy_task = tokio::spawn(async {});

tokio::time::timeout(
Duration::from_secs(5),
reader.disconnect(Error::Socket("boom".to_string()), &mut dummy_task),
reader.disconnect(Error::Socket("boom".to_string())),
)
.await
.expect("disconnect must not block on a stalled stream");
Expand All @@ -605,3 +626,7 @@ mod teardown_tests {
);
}
}

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