From b6f10c0962a5a3853c63ecb9b800256d7d1b157a Mon Sep 17 00:00:00 2001 From: Tim Zhang Date: Mon, 17 Aug 2026 16:20:00 +0800 Subject: [PATCH] async/client: Fix unary request timeout cleanup Apply one absolute deadline while a unary request waits for outbound queue capacity, socket writes, and its response. Remove stream registrations when requests time out or are cancelled. Discard untransformed requests that expire before writing. Once a stateful transform advances its state, keep the frame queued to avoid desynchronizing the connection. Close the connection when timeout or cancellation interrupts an in-progress write because the frame may be incomplete. Return writer failures through the task result and add regression tests for queue saturation, response cleanup, queued request timeouts, and cancellation during writes. Use synchronous fast paths for channel capacity, outbound locking, and ordinary writes. Only construct timeout futures when contention requires waiting. Part of #317 Signed-off-by: Tim Zhang --- example-prost/async_client.rs | 16 +- example/async-client.rs | 16 +- src/asynchronous/client.rs | 123 ++++++----- src/asynchronous/client_tests.rs | 297 +++++++++++++++++++++++++++ src/asynchronous/connection.rs | 172 +++++++++++----- src/asynchronous/mod.rs | 3 +- src/asynchronous/server.rs | 3 +- src/asynchronous/stream.rs | 40 ++++ src/security_extension.rs | 141 +++++++++++-- tests/hook_integration_async_unix.rs | 3 +- 10 files changed, 677 insertions(+), 137 deletions(-) create mode 100644 src/asynchronous/client_tests.rs diff --git a/example-prost/async_client.rs b/example-prost/async_client.rs index ff4e3040..1aadfd4e 100644 --- a/example-prost/async_client.rs +++ b/example-prost/async_client.rs @@ -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: {:?}", diff --git a/example/async-client.rs b/example/async-client.rs index 39f0439d..4567eb46 100644 --- a/example/async-client.rs +++ b/example/async-client.rs @@ -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!( diff --git a/src/asynchronous/client.rs b/src/asynchronous/client.rs index 2421b6e4..cc4bf856 100644 --- a/src/asynchronous/client.rs +++ b/src/asynchronous/client.rs @@ -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::{ @@ -21,7 +25,6 @@ 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, }; @@ -29,9 +32,37 @@ use crate::r#async::stream::{ 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>, + 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 @@ -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 { 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)? @@ -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?; @@ -281,18 +316,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 + 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, @@ -300,7 +335,9 @@ impl Client { streaming_server, self.streams.clone(), self.conn_ctx.clone(), - )) + ); + registration.disarm(); + Ok(inner) } } @@ -316,16 +353,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, }, ) } @@ -333,7 +367,6 @@ impl Builder for ClientBuilder { struct ClientWriter { rx: MessageReceiver, - shutdown_notifier: shutdown::Notifier, } #[async_trait] @@ -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, } #[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 @@ -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}; @@ -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"); @@ -605,3 +626,7 @@ mod teardown_tests { ); } } + +#[cfg(test)] +#[path = "client_tests.rs"] +mod tests; diff --git a/src/asynchronous/client_tests.rs b/src/asynchronous/client_tests.rs new file mode 100644 index 00000000..ba04c13a --- /dev/null +++ b/src/asynchronous/client_tests.rs @@ -0,0 +1,297 @@ +use std::pin::Pin; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::sync::Arc; +use std::task::{Context, Poll}; +use std::time::Duration; + +use futures::future::join_all; +use futures::task::AtomicWaker; +use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; +use tokio::sync::Notify; +use tokio::time::timeout; + +use super::*; + +struct BlockedSocket { + write_polls: Arc, + write_started: Arc, +} + +struct SinkSocket; + +struct WriteGate { + open: AtomicBool, + bytes_written: AtomicUsize, + waker: AtomicWaker, +} + +impl WriteGate { + fn new() -> Self { + Self { + open: AtomicBool::new(false), + bytes_written: AtomicUsize::new(0), + waker: AtomicWaker::new(), + } + } + + fn open(&self) { + self.open.store(true, Ordering::Release); + self.waker.wake(); + } + + fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<()> { + if self.open.load(Ordering::Acquire) { + return Poll::Ready(()); + } + self.waker.register(cx.waker()); + if self.open.load(Ordering::Acquire) { + Poll::Ready(()) + } else { + Poll::Pending + } + } +} + +struct GatedSocket { + gate: Arc, +} + +impl AsyncRead for BlockedSocket { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + _buf: &mut ReadBuf<'_>, + ) -> Poll> { + Poll::Pending + } +} + +impl AsyncWrite for BlockedSocket { + fn poll_write( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + _buf: &[u8], + ) -> Poll> { + self.write_polls.fetch_add(1, Ordering::Relaxed); + self.write_started.notify_waiters(); + Poll::Pending + } + + fn poll_flush( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Pending + } + + fn poll_shutdown( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } +} + +impl AsyncRead for SinkSocket { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + _buf: &mut ReadBuf<'_>, + ) -> Poll> { + Poll::Pending + } +} + +impl AsyncWrite for SinkSocket { + fn poll_write( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_shutdown( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } +} + +impl AsyncRead for GatedSocket { + fn poll_read( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + _buf: &mut ReadBuf<'_>, + ) -> Poll> { + Poll::Pending + } +} + +impl AsyncWrite for GatedSocket { + fn poll_write( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + match self.gate.poll_ready(cx) { + Poll::Ready(()) => { + self.gate + .bytes_written + .fetch_add(buf.len(), Ordering::Relaxed); + Poll::Ready(Ok(buf.len())) + } + Poll::Pending => Poll::Pending, + } + } + + fn poll_flush( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll> { + self.gate.poll_ready(cx).map(|()| Ok(())) + } + + fn poll_shutdown( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + ) -> Poll> { + Poll::Ready(Ok(())) + } +} + +fn blocked_client() -> (Client, Arc, Arc) { + let write_polls = Arc::new(AtomicUsize::new(0)); + let write_started = Arc::new(Notify::new()); + let socket = BlockedSocket { + write_polls: write_polls.clone(), + write_started: write_started.clone(), + }; + ( + Client::new(Socket::new(socket)), + write_polls, + write_started, + ) +} + +fn request_with_timeout(timeout: Duration) -> Request { + Request { + timeout_nano: timeout.as_nanos() as i64, + ..Default::default() + } +} + +#[tokio::test] +async fn request_deadline_covers_a_full_outbound_queue() { + let (client, write_polls, _) = blocked_client(); + let mut tasks = Vec::new(); + + for _ in 0..110 { + let client = client.clone(); + tasks.push(tokio::spawn(async move { + client + .request(request_with_timeout(Duration::from_millis(100))) + .await + })); + } + + let results = timeout(Duration::from_secs(2), join_all(tasks)) + .await + .expect("requests must not remain blocked behind the full queue"); + assert!(results + .into_iter() + .all(|result| result.expect("request task panicked").is_err())); + assert!(write_polls.load(Ordering::Relaxed) > 0); + assert!(client.streams.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn response_timeout_removes_the_stream_without_closing_the_connection() { + let client = Client::new(Socket::new(SinkSocket)); + + let result = client + .request(request_with_timeout(Duration::from_millis(50))) + .await; + + assert!(result.is_err()); + assert!(client.streams.lock().unwrap().is_empty()); + assert!(!client.req_tx.is_closed()); +} + +#[tokio::test] +async fn expired_queued_request_preserves_the_timeout_error() { + let gate = Arc::new(WriteGate::new()); + let client = Client::new(Socket::new(GatedSocket { gate: gate.clone() })); + let blocker = GenMessage { + header: MessageHeader::new_data(2, 0), + payload: Vec::new(), + }; + client + .req_tx + .send(SendingMessage::new(blocker)) + .await + .unwrap(); + tokio::task::yield_now().await; + + let request = client.request(request_with_timeout(Duration::from_millis(50))); + tokio::pin!(request); + assert!(futures::poll!(request.as_mut()).is_pending()); + let marker = GenMessage { + header: MessageHeader::new_data(4, 0), + payload: Vec::new(), + }; + client + .req_tx + .send(SendingMessage::new(marker)) + .await + .unwrap(); + + tokio::time::sleep(Duration::from_millis(75)).await; + gate.open(); + timeout(Duration::from_secs(1), async { + while gate.bytes_written.load(Ordering::Relaxed) + < 2 * crate::proto::MESSAGE_HEADER_LENGTH + { + tokio::task::yield_now().await; + } + }) + .await + .expect("writer did not process the message after the expired request"); + + assert_eq!( + request.await, + Err(Error::Others("Request deadline elapsed".to_string())) + ); + assert_eq!( + gate.bytes_written.load(Ordering::Relaxed), + 2 * crate::proto::MESSAGE_HEADER_LENGTH + ); + assert!(client.streams.lock().unwrap().is_empty()); + assert!(!client.req_tx.is_closed()); +} + +#[tokio::test] +async fn cancelling_an_in_progress_write_cleans_up_and_closes_the_connection() { + let (client, _, write_started) = blocked_client(); + let request_client = client.clone(); + let request = tokio::spawn(async move { request_client.request(Request::default()).await }); + + timeout(Duration::from_secs(1), write_started.notified()) + .await + .expect("writer did not start"); + request.abort(); + request.await.expect_err("request task was not cancelled"); + + timeout(Duration::from_secs(1), client.req_tx.closed()) + .await + .expect("connection was not closed after cancelling an in-progress write"); + assert!(client.streams.lock().unwrap().is_empty()); +} diff --git a/src/asynchronous/connection.rs b/src/asynchronous/connection.rs index ef633c26..ffd3d8c2 100644 --- a/src/asynchronous/connection.rs +++ b/src/asynchronous/connection.rs @@ -6,15 +6,107 @@ use async_trait::async_trait; use log::{error, trace}; -use tokio::io::split; -use tokio::sync::oneshot; +use tokio::io::{split, AsyncWrite}; +use tokio::time::{sleep_until, Instant}; use tokio::{io::ReadHalf, select, task}; -use crate::error::Error; +use crate::error::{Error, Result}; use crate::proto::{GenMessage, GenMessageError, MessageHeader}; use super::{stream::SendingMessage, transport::Socket}; +enum WriteOutcome { + Complete(crate::error::Result<()>), + Discarded(Error), + Cancelled, + DeadlineElapsed, +} + +pub(crate) fn request_timeout_error() -> Error { + Error::Others("Request deadline elapsed".to_string()) +} + +async fn write_message( + writer: &mut (impl AsyncWrite + Unpin), + sending_msg: &SendingMessage, +) -> WriteOutcome { + let Some(control) = sending_msg.control.as_ref() else { + trace!("write message: {:?}", sending_msg.msg); + return WriteOutcome::Complete(sending_msg.msg.write_to(writer).await); + }; + let deadline = control.deadline(); + + let expired = deadline.is_some_and(|deadline| deadline <= Instant::now()); + if control.is_cancelled() { + return WriteOutcome::Discarded(Error::LocalClosed); + } + if expired { + return WriteOutcome::Discarded(request_timeout_error()); + } + + trace!("write message: {:?}", sending_msg.msg); + if let Some(deadline) = deadline { + select! { + biased; + result = sending_msg.msg.write_to(writer) => WriteOutcome::Complete(result), + _ = control.cancelled() => WriteOutcome::Cancelled, + _ = sleep_until(deadline) => WriteOutcome::DeadlineElapsed, + } + } else { + select! { + biased; + result = sending_msg.msg.write_to(writer) => WriteOutcome::Complete(result), + _ = control.cancelled() => WriteOutcome::Cancelled, + } + } +} + +async fn run_writer( + mut writer: impl AsyncWrite + Unpin, + mut writer_delegate: impl WriterDelegate, +) -> Result<()> { + let result = loop { + let Some(mut sending_msg) = writer_delegate.recv().await else { + break Ok(()); + }; + + let failure = match write_message(&mut writer, &sending_msg).await { + WriteOutcome::Complete(Ok(())) => { + sending_msg.send_result(Ok(())); + continue; + } + WriteOutcome::Discarded(e) => { + sending_msg.send_result(Err(e)); + continue; + } + WriteOutcome::Complete(Err(e)) => Some((e.clone(), e)), + WriteOutcome::Cancelled => Some(( + Error::LocalClosed, + Error::Socket( + "connection closed after a request was cancelled during write".to_string(), + ), + )), + WriteOutcome::DeadlineElapsed => Some(( + request_timeout_error(), + Error::Socket( + "connection closed after a request deadline elapsed during write".to_string(), + ), + )), + }; + + if let Some((message_error, connection_error)) = failure { + error!("write_message got error: {:?}", connection_error); + sending_msg.send_result(Err(message_error)); + // Return without waiting for socket shutdown, which may never complete. + break Err(connection_error); + } + }; + + writer_delegate.exit().await; + trace!("Writer task exit."); + result +} + pub trait Builder { type Reader; type Writer; @@ -31,7 +123,7 @@ pub trait WriterDelegate { #[async_trait] pub trait ReaderDelegate { async fn wait_shutdown(&self); - async fn disconnect(&self, e: Error, task: &mut task::JoinHandle<()>); + async fn disconnect(&self, e: Error); async fn exit(&self); async fn handle_msg(&self, msg: GenMessage); async fn handle_err(&self, header: MessageHeader, e: Error); @@ -39,13 +131,8 @@ pub trait ReaderDelegate { pub struct Connection { reader: ReadHalf, - writer_task: task::JoinHandle<()>, + writer_task: task::JoinHandle>, reader_delegate: B::Reader, - // Delivers a fatal write error from the writer task. Receiving a value - // means the writer hit an unrecoverable transport error and the whole - // connection must be torn down; the channel closing without a value means - // the writer stopped normally. - writer_error: oneshot::Receiver, } impl Connection @@ -55,45 +142,17 @@ where B::Writer: WriterDelegate + Send + Sync + 'static, { pub fn new(conn: Socket, mut builder: B) -> Self { - let (reader, mut writer) = split(conn); + let (reader, writer) = split(conn); - let (reader_delegate, mut writer_delegate) = builder.build(); - let (err_tx, err_rx) = oneshot::channel(); + let (reader_delegate, writer_delegate) = builder.build(); // Long-running sender task - let writer_task = tokio::spawn(async move { - while let Some(mut sending_msg) = writer_delegate.recv().await { - trace!("write message: {:?}", sending_msg.msg); - if let Err(e) = sending_msg.msg.write_to(&mut writer).await { - error!("write_message got error: {:?}", e); - // Report the failure to the caller awaiting this send. - sending_msg.send_result(Err(e.clone())); - // write_to uses write_all internally, so a failed write may - // have left a partial frame on the wire, desynchronizing the - // frame boundaries of every stream multiplexed on this - // connection; it is no longer usable. (This can happen on an - // otherwise healthy socket, e.g. ENOMEM when the kernel - // cannot satisfy a high-order allocation under memory - // fragmentation.) Report the error to Connection::run at once - // and exit. Deliberately do NOT wait on writer.shutdown(): it - // can block (some transports never complete poll_shutdown) - // and would delay or prevent cleanup. run() closes the whole - // connection; dropping this task drops the write half. - let _ = err_tx.send(e); - return; - } - sending_msg.send_result(Ok(())); - } - // The outbound channel closed: this is a normal shutdown. - writer_delegate.exit().await; - trace!("Writer task exit."); - }); + let writer_task = tokio::spawn(run_writer(writer, writer_delegate)); Self { reader, writer_task, reader_delegate, - writer_error: err_rx, } } @@ -102,24 +161,23 @@ where mut reader, mut writer_task, reader_delegate, - mut writer_error, } = self; loop { select! { - // Fixed poll order: a write error wins over the shutdown - // notification the writer raises as it unwinds, and a pending - // shutdown wins over further reads. Both are idle during normal - // operation, so read_from is still reached every iteration. + // Writer failures take priority, then shutdown, then incoming frames. biased; - - werr = &mut writer_error => { - // Ok(e): the writer hit a fatal transport error. Drive a - // connection-wide teardown — fail every registered client - // stream / stop the server handlers, and drop the read half - // on exit. Err(_): the writer stopped without an error. - if let Ok(e) = werr { - trace!("Writer failed, tearing down connection: {:?}", e); - reader_delegate.disconnect(e, &mut writer_task).await; + writer_result = &mut writer_task => { + match writer_result { + Ok(Ok(())) => {} + Ok(Err(e)) => { + trace!("Write msg err: {:?}", e); + reader_delegate.disconnect(e).await; + } + Err(e) => { + let e = Error::Others(format!("Writer task failed: {e}")); + error!("Write task err: {:?}", e); + reader_delegate.disconnect(e).await; + } } break; } @@ -140,7 +198,9 @@ where Err(GenMessageError::InternalError(e)) => { trace!("Read msg err: {:?}", e); - reader_delegate.disconnect(e, &mut writer_task).await; + writer_task.abort(); + let _ = (&mut writer_task).await; + reader_delegate.disconnect(e).await; break; } } diff --git a/src/asynchronous/mod.rs b/src/asynchronous/mod.rs index 5a7e549e..bcb0848f 100644 --- a/src/asynchronous/mod.rs +++ b/src/asynchronous/mod.rs @@ -32,7 +32,8 @@ pub use self::stream::{ SSSender, ServerStream, ServerStreamReceiver, ServerStreamSender, StreamInner, StreamReceiver, StreamSender, }; -pub(crate) use self::stream::SendingMessage; +pub(crate) use self::stream::{MessageControl, SendingMessage}; +pub(crate) use connection::request_timeout_error; #[doc(inline)] pub use crate::r#async::client::Client; #[doc(inline)] diff --git a/src/asynchronous/server.rs b/src/asynchronous/server.rs index b7e00b12..3f0bf463 100644 --- a/src/asynchronous/server.rs +++ b/src/asynchronous/server.rs @@ -17,7 +17,6 @@ use futures::StreamExt as _; use tokio::{ self, select, spawn, sync::mpsc::{channel, Sender}, - task, time::timeout, }; @@ -435,7 +434,7 @@ impl ReaderDelegate for ServerReader { self.server_shutdown.wait_shutdown().await } - async fn disconnect(&self, _: Error, _: &mut task::JoinHandle<()>) { + async fn disconnect(&self, _: Error) { self.handler_shutdown.shutdown(); // TODO: Don't wait for all requests to complete? when the connection is disconnected. } diff --git a/src/asynchronous/stream.rs b/src/asynchronous/stream.rs index 2b10cc7e..1bed4a22 100644 --- a/src/asynchronous/stream.rs +++ b/src/asynchronous/stream.rs @@ -12,6 +12,7 @@ use std::sync::{Arc, Mutex}; use crate::ConnectionContext; use tokio::sync::mpsc; +use tokio::time::Instant; use super::Client; use crate::error::{Error, Result}; @@ -71,6 +72,34 @@ impl StreamRegistry { pub struct SendingMessage { pub msg: GenMessage, pub result_chan: Option>>, + pub(crate) control: Option, +} + +#[derive(Debug)] +pub(crate) struct MessageControl { + deadline: Option, + response_tx: ClientResultSender, +} + +impl MessageControl { + pub(crate) fn new(deadline: Option, response_tx: ClientResultSender) -> Self { + Self { + deadline, + response_tx, + } + } + + pub(crate) fn deadline(&self) -> Option { + self.deadline + } + + pub(crate) fn is_cancelled(&self) -> bool { + self.response_tx.is_closed() + } + + pub(crate) async fn cancelled(&self) { + self.response_tx.closed().await; + } } impl SendingMessage { @@ -78,8 +107,18 @@ impl SendingMessage { Self { msg, result_chan: None, + control: None, + } + } + + pub(crate) fn new_with_control(msg: GenMessage, control: MessageControl) -> Self { + Self { + msg, + result_chan: None, + control: Some(control), } } + pub fn new_with_result( msg: GenMessage, result_chan: tokio::sync::oneshot::Sender>, @@ -87,6 +126,7 @@ impl SendingMessage { Self { msg, result_chan: Some(result_chan), + control: None, } } diff --git a/src/security_extension.rs b/src/security_extension.rs index ccb8985a..643e8971 100644 --- a/src/security_extension.rs +++ b/src/security_extension.rs @@ -188,6 +188,27 @@ pub(crate) use hooks::ServerExtensionConfig; #[cfg(feature = "security_extension")] pub use hooks::{AcceptHook, ConnectHook, HookError, HookOutput}; +#[cfg(feature = "async")] +async fn reserve_message_slot<'a>( + tx: &'a tokio::sync::mpsc::Sender, + deadline: Option, +) -> Result, Error> { + match tx.try_reserve() { + Ok(permit) => Ok(permit), + Err(_) => { + let reserve = tx.reserve(); + if let Some(deadline) = deadline { + tokio::time::timeout_at(deadline, reserve) + .await + .map_err(|_| crate::asynchronous::request_timeout_error())? + } else { + reserve.await + } + .map_err(|e| Error::Others(format!("reserve channel capacity failed: {e}"))) + } + } +} + // ── Feature-gated hook types ─────────────────────────────────────────────── // // All hook-related items live in this inner module behind a single cfg gate. @@ -703,30 +724,83 @@ mod hooks { rpc_error: bool, await_ack: bool, ) -> Result<(), Error> { - // Reserve capacity first — this is the only cancellable await point. - // If cancelled here, no nonce has been advanced. - let permit = tx - .reserve() + self.transform_send_inner(msg, tx, rpc_error, await_ack, None) + .await + } + + #[cfg(feature = "async")] + pub(crate) async fn transform_send_with_control( + &self, + msg: &mut crate::proto::GenMessage, + tx: &tokio::sync::mpsc::Sender, + rpc_error: bool, + await_ack: bool, + control: crate::asynchronous::MessageControl, + ) -> Result<(), Error> { + self.transform_send_inner(msg, tx, rpc_error, await_ack, Some(control)) .await - .map_err(|e| Error::Others(format!("reserve channel capacity failed: {e}")))?; + } - let _guard = self.async_outbound_lock.lock().await; + #[cfg(feature = "async")] + async fn transform_send_inner( + &self, + msg: &mut crate::proto::GenMessage, + tx: &tokio::sync::mpsc::Sender, + rpc_error: bool, + await_ack: bool, + control: Option, + ) -> Result<(), Error> { + let deadline = control.as_ref().and_then(|control| control.deadline()); + // Reserve capacity first — this is the only cancellable await point. + // If cancelled here, no nonce has been advanced. + let permit = reserve_message_slot(tx, deadline).await?; + + let _guard = match self.async_outbound_lock.try_lock() { + Ok(guard) => guard, + Err(_) => { + let lock = self.async_outbound_lock.lock(); + if let Some(deadline) = deadline { + tokio::time::timeout_at(deadline, lock) + .await + .map_err(|_| crate::asynchronous::request_timeout_error())? + } else { + lock.await + } + } + }; self.outbound(msg, rpc_error)?; let taken = std::mem::take(msg); + // A stateful transform may advance a nonce or counter. Once that + // happens, the frame must not be discarded by the writer on + // timeout or cancellation, or the peers will become desynchronized. + // The caller's deadline still bounds reserve and lock acquisition. + let control = if self.payload_transform.is_some() { + None + } else { + control + }; + // From here on: no await until the frame is in the channel. // permit.send() is synchronous — cannot be cancelled. if await_ack { let (result_tx, result_rx) = tokio::sync::oneshot::channel(); - permit.send(crate::asynchronous::SendingMessage::new_with_result( - taken, result_tx, - )); + let mut sending_msg = + crate::asynchronous::SendingMessage::new_with_result(taken, result_tx); + sending_msg.control = control; + permit.send(sending_msg); drop(_guard); result_rx .await .map_err(|_| Error::Others("writer task dropped result channel".to_string()))? } else { - permit.send(crate::asynchronous::SendingMessage::new(taken)); + let sending_msg = match control { + Some(control) => { + crate::asynchronous::SendingMessage::new_with_control(taken, control) + } + None => crate::asynchronous::SendingMessage::new(taken), + }; + permit.send(sending_msg); Ok(()) } } @@ -847,22 +921,53 @@ mod hooks { rpc_error: bool, await_ack: bool, ) -> Result<(), Error> { - self.outbound(msg, rpc_error)?; - let permit = tx - .reserve() + self.transform_send_inner(msg, tx, rpc_error, await_ack, None) + .await + } + + #[cfg(feature = "async")] + pub(crate) async fn transform_send_with_control( + &self, + msg: &mut crate::proto::GenMessage, + tx: &tokio::sync::mpsc::Sender, + rpc_error: bool, + await_ack: bool, + control: crate::asynchronous::MessageControl, + ) -> Result<(), Error> { + self.transform_send_inner(msg, tx, rpc_error, await_ack, Some(control)) .await - .map_err(|e| Error::Others(format!("reserve channel capacity failed: {e}")))?; + } + + #[cfg(feature = "async")] + async fn transform_send_inner( + &self, + msg: &mut crate::proto::GenMessage, + tx: &tokio::sync::mpsc::Sender, + rpc_error: bool, + await_ack: bool, + control: Option, + ) -> Result<(), Error> { + self.outbound(msg, rpc_error)?; + let deadline = control.as_ref().and_then(|control| control.deadline()); + let permit = reserve_message_slot(tx, deadline).await?; let taken = std::mem::take(msg); if await_ack { let (result_tx, result_rx) = tokio::sync::oneshot::channel(); - permit.send(crate::asynchronous::SendingMessage::new_with_result( - taken, result_tx, - )); + let mut sending_msg = + crate::asynchronous::SendingMessage::new_with_result(taken, result_tx); + sending_msg.control = control; + permit.send(sending_msg); result_rx .await .map_err(|_| Error::Others("writer task dropped result channel".to_string()))? } else { - permit.send(crate::asynchronous::SendingMessage::new(taken)); + let sending_msg = match control { + Some(control) => { + crate::asynchronous::SendingMessage::new_with_control(taken, control) + } + None => crate::asynchronous::SendingMessage::new(taken), + }; + permit.send(sending_msg); Ok(()) } } diff --git a/tests/hook_integration_async_unix.rs b/tests/hook_integration_async_unix.rs index f6c7eda1..3499ac09 100644 --- a/tests/hook_integration_async_unix.rs +++ b/tests/hook_integration_async_unix.rs @@ -1072,7 +1072,7 @@ async fn test_server_initiated_stream_close_client_gets_final_response() { // Test 6: Unary request timeout (server-side DEADLINE_EXCEEDED + client-side timeout) // // Covers the timeout code path in server.rs handle_method() (tokio::time::timeout -// around the handler) and client.rs request() (tokio::time::timeout on the response). +// around the handler) and client.rs request() (one deadline for send and response). #[tokio::test] async fn test_unary_request_timeout() { let sock_path = temp_unix_socket_path(); @@ -1104,6 +1104,7 @@ async fn test_unary_request_timeout() { assert!( err_str.contains("timeout") || err_str.contains("Timeout") + || err_str.contains("deadline elapsed") || err_str.contains("DEADLINE_EXCEEDED"), "Expected timeout-related error, got: {}", err_str