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
4 changes: 1 addition & 3 deletions example/async-client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -39,9 +39,7 @@ async fn main() {

assert_eq!(
resp,
Err(ttrpc::Error::Others(
"Receive packet timeout Elapsed(())".into()
))
Err(ttrpc::Error::Others("Request deadline elapsed".into()))
);

println!(
Expand Down
112 changes: 74 additions & 38 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::ConnectionContext;
Expand All @@ -23,14 +27,41 @@ use crate::proto::{
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>>,
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 ttrpc Client (async).
#[derive(Clone)]
pub struct Client {
Expand Down Expand Up @@ -155,35 +186,43 @@ impl Client {
/// Requests 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 mut 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(),
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 @@ -229,14 +268,18 @@ 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.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(&mut msg, &self.req_tx, false, false)
.await?;

Ok(StreamInner::new(
let inner = StreamInner::new(
stream_id,
self.req_tx.clone(),
rx,
Expand All @@ -245,7 +288,9 @@ impl Client {
Kind::Client,
self.streams.clone(),
self.conn_ctx.clone(),
))
);
registration.disarm();
Ok(inner)
}
}

Expand All @@ -261,16 +306,13 @@ 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,
streams: self.streams.clone(),
},
)
Expand All @@ -279,8 +321,6 @@ impl Builder for ClientBuilder {

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

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

Expand Down Expand Up @@ -308,9 +348,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 @@ -367,22 +405,16 @@ async fn get_resp_tx(

struct ClientReader {
streams: Arc<Mutex<HashMap<u32, ResultSender>>>,
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 undone RPC requests with the error.
Expand Down Expand Up @@ -430,7 +462,7 @@ impl ReaderDelegate for 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 @@ -468,3 +500,7 @@ mod tests {
);
}
}

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