1use std::sync::Arc;
5
6use futures::{SinkExt, StreamExt};
7use tokio::io::{AsyncReadExt, ReadHalf, WriteHalf};
8use tokio::{
9 io::AsyncWriteExt,
10 net::TcpStream,
11 time::{self, Duration, Instant},
12};
13use tokio_util::codec::{FramedRead, FramedWrite};
14
15use prometheus::IntCounter;
16
17use super::{CallHomeHandshake, ControlMessage, TcpStreamConnectionInfo};
18use crate::engine::AsyncEngineContext;
19use crate::pipeline::network::{
20 ConnectionInfo, ResponseStreamPrologue, StreamReceiver, StreamSender,
21 codec::{TwoPartCodec, TwoPartMessage, TwoPartMessageType},
22 tcp::StreamType,
23};
24use anyhow::{Context, Result, anyhow as error}; #[allow(dead_code)]
27pub struct TcpClient {
28 worker_id: String,
29}
30
31impl Default for TcpClient {
32 fn default() -> Self {
33 TcpClient {
34 worker_id: uuid::Uuid::new_v4().to_string(),
35 }
36 }
37}
38
39impl TcpClient {
40 pub fn new(worker_id: String) -> Self {
41 TcpClient { worker_id }
42 }
43
44 async fn connect(address: &str) -> std::io::Result<TcpStream> {
45 let backoff = std::time::Duration::from_millis(200);
47 loop {
48 match TcpStream::connect(address).await {
49 Ok(socket) => {
50 socket.set_nodelay(true)?;
51 return Ok(socket);
52 }
53 Err(e) => {
54 if e.kind() == std::io::ErrorKind::AddrNotAvailable {
55 tracing::warn!("retry warning: failed to connect: {:?}", e);
56 tokio::time::sleep(backoff).await;
57 } else {
58 return Err(e);
59 }
60 }
61 }
62 }
63 }
64
65 pub async fn create_response_stream(
66 context: Arc<dyn AsyncEngineContext>,
67 info: ConnectionInfo,
68 cancellation_counter: Option<IntCounter>,
69 ) -> Result<StreamSender> {
70 let info =
71 TcpStreamConnectionInfo::try_from(info).context("tcp-stream-connection-info-error")?;
72 tracing::trace!("Creating response stream for {:?}", info);
73
74 if info.stream_type != StreamType::Response {
75 return Err(error!(
76 "Invalid stream type; TcpClient requires the stream type to be `response`; however {:?} was passed",
77 info.stream_type
78 ));
79 }
80
81 if info.context != context.id() {
82 return Err(error!(
83 "Invalid context; TcpClient requires the context to be {:?}; however {:?} was passed",
84 context.id(),
85 info.context
86 ));
87 }
88
89 let stream = TcpClient::connect(&info.address).await?;
90 let peer_port = stream.peer_addr().ok().map(|addr| addr.port());
91 let (read_half, write_half) = tokio::io::split(stream);
92
93 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
94 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
95
96 let (alive_tx, alive_rx) = tokio::sync::oneshot::channel::<()>();
102
103 let reader_task = tokio::spawn(handle_reader(
104 framed_reader,
105 context.clone(),
106 alive_tx,
107 cancellation_counter,
108 ));
109
110 let handshake = CallHomeHandshake {
112 subject: info.subject.clone(),
113 stream_type: StreamType::Response,
114 };
115
116 let handshake_bytes = match serde_json::to_vec(&handshake) {
117 Ok(hb) => hb,
118 Err(err) => {
119 return Err(error!(
120 "create_response_stream: Error converting CallHomeHandshake to JSON array: {err:#}"
121 ));
122 }
123 };
124 let msg = TwoPartMessage::from_header(handshake_bytes.into());
125
126 framed_writer
128 .send(msg)
129 .await
130 .map_err(|e| error!("failed to send handshake: {:?}", e))?;
131
132 let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel(64);
134
135 let writer_context = context.clone();
137 let writer_task = tokio::spawn(handle_writer(
138 framed_writer,
139 bytes_rx,
140 alive_rx,
141 writer_context,
142 ));
143
144 let subject = info.subject.clone();
145 let monitor_context = context;
146 tokio::spawn(async move {
149 let _ = wait_for_connection_tasks(
150 reader_task,
151 writer_task,
152 monitor_context,
153 peer_port,
154 subject,
155 )
156 .await;
157 });
158
159 let prologue = Some(ResponseStreamPrologue { error: None });
162
163 let stream_sender = StreamSender {
165 tx: bytes_tx,
166 prologue,
167 };
168
169 Ok(stream_sender)
170 }
171
172 pub async fn create_request_stream(
186 context: Arc<dyn AsyncEngineContext>,
187 info: ConnectionInfo,
188 cancellation_counter: Option<IntCounter>,
189 ) -> Result<StreamReceiver> {
190 let info =
191 TcpStreamConnectionInfo::try_from(info).context("tcp-stream-connection-info-error")?;
192 tracing::trace!("Creating request stream for {:?}", info);
193
194 if info.stream_type != StreamType::Request {
195 return Err(error!(
196 "Invalid stream type; TcpClient::create_request_stream requires the stream type to be `request`; however {:?} was passed",
197 info.stream_type
198 ));
199 }
200
201 if info.context != context.id() {
202 return Err(error!(
203 "Invalid context; TcpClient::create_request_stream requires the context to be {:?}; however {:?} was passed",
204 context.id(),
205 info.context
206 ));
207 }
208
209 let stream = TcpClient::connect(&info.address).await?;
210 let (read_half, write_half) = tokio::io::split(stream);
211
212 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
213 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
214
215 let handshake = CallHomeHandshake {
216 subject: info.subject.clone(),
217 stream_type: StreamType::Request,
218 };
219 let handshake_bytes = serde_json::to_vec(&handshake).map_err(|err| {
220 error!(
221 "create_request_stream: Error converting CallHomeHandshake to JSON array: {err:#}"
222 )
223 })?;
224 framed_writer
225 .send(TwoPartMessage::from_header(handshake_bytes.into()))
226 .await
227 .map_err(|e| error!("failed to send request-stream handshake: {:?}", e))?;
228
229 drop(framed_writer);
232
233 let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel::<bytes::Bytes>(64);
234
235 tokio::spawn(handle_request_reader(
236 framed_reader,
237 bytes_tx,
238 context,
239 cancellation_counter,
240 ));
241
242 Ok(StreamReceiver { rx: bytes_rx })
243 }
244}
245
246async fn handle_request_reader(
247 mut framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
248 bytes_tx: tokio::sync::mpsc::Sender<bytes::Bytes>,
249 context: Arc<dyn AsyncEngineContext>,
250 cancellation_counter: Option<IntCounter>,
251) {
252 let cancellation_seen = {
253 let killed = context.killed();
256 let stopped = context.stopped();
257 let bytes_closed_tx = bytes_tx.clone();
260 let bytes_closed = bytes_closed_tx.closed();
261 tokio::pin!(killed, stopped, bytes_closed);
262
263 let mut cancellation_seen = false;
265 loop {
266 tokio::select! {
267 biased;
268
269 _ = &mut killed => {
270 tracing::trace!("context kill signal received on request stream; shutting down");
271 break;
272 }
273
274 _ = &mut stopped => {
275 tracing::trace!("context stop signal received on request stream; shutting down");
276 break;
277 }
278
279 _ = &mut bytes_closed => {
284 tracing::debug!("downstream consumer dropped; exiting request-stream reader");
285 break;
286 }
287
288 msg = framed_reader.next() => {
289 match msg {
290 Some(Ok(two_part_msg)) => match two_part_msg.into_message_type() {
291 TwoPartMessageType::HeaderOnly(header) => {
292 let ctrl = match serde_json::from_slice::<ControlMessage>(&header) {
293 Ok(c) => c,
294 Err(e) => {
295 tracing::warn!(
296 err = ?e,
297 "invalid control message, closing connection"
298 );
299 cancellation_seen = true;
300 context.kill();
301 break;
302 }
303 };
304 match ctrl {
305 ControlMessage::Stop => {
306 cancellation_seen = true;
307 context.stop();
308 break;
309 }
310 ControlMessage::Kill => {
311 cancellation_seen = true;
312 context.kill();
313 break;
314 }
315 ControlMessage::Sentinel => {
316 tracing::trace!("upstream signaled end of request stream");
317 break;
318 }
319 }
320 }
321 TwoPartMessageType::DataOnly(data) => {
322 if bytes_tx.send(data).await.is_err() {
323 tracing::debug!("downstream consumer dropped; exiting request-stream reader");
324 break;
325 }
326 }
327 _ => {
328 tracing::warn!("fatal error - unexpected message shape on request stream");
329 cancellation_seen = true;
330 context.kill();
331 break;
332 }
333 }
334 Some(Err(e)) => {
335 tracing::warn!("fatal error - failed to decode message on request stream: {e:?}");
336 cancellation_seen = true;
337 context.kill();
338 break;
339 }
340 None => {
341 tracing::warn!("request stream closed by upstream before sentinel; treating as truncated");
346 cancellation_seen = true;
347 context.kill();
348 break;
349 }
350 }
351 }
352 }
353 }
354
355 cancellation_seen
358 };
359
360 if cancellation_seen && let Some(counter) = &cancellation_counter {
361 counter.inc();
362 }
363
364 drop(bytes_tx);
367}
368
369async fn wait_for_connection_tasks(
370 reader_task: tokio::task::JoinHandle<FramedRead<ReadHalf<TcpStream>, TwoPartCodec>>,
371 writer_task: tokio::task::JoinHandle<Result<FramedWrite<WriteHalf<TcpStream>, TwoPartCodec>>>,
372 context: Arc<dyn AsyncEngineContext>,
373 peer_port: Option<u16>,
374 subject: String,
375) -> Result<()> {
376 let reader = match reader_task.await {
379 Ok(reader) => reader,
380 Err(reader_err) => {
381 writer_task.abort();
382 let _ = writer_task.await;
383 tracing::error!(
384 subject = %subject,
385 peer_port = ?peer_port,
386 err = ?reader_err,
387 "reader task failed to join"
388 );
389 return Err(reader_err.into());
390 }
391 };
392
393 let writer = match writer_task.await {
394 Ok(writer) => writer,
395 Err(writer_err) => {
396 tracing::error!(
397 subject = %subject,
398 peer_port = ?peer_port,
399 err = ?writer_err,
400 "writer task failed to join"
401 );
402 return Err(writer_err.into());
403 }
404 };
405
406 let reader = reader.into_inner();
407 let writer = match writer {
408 Ok(writer) => writer.into_inner(),
409 Err(e) => {
410 tracing::error!(
411 subject = %subject,
412 peer_port = ?peer_port,
413 err = ?e,
414 "writer task returned error"
415 );
416 return Err(e);
417 }
418 };
419
420 let stream = reader.unsplit(writer);
421 wait_for_server_shutdown(stream, context).await
422}
423
424async fn wait_for_server_shutdown(
425 mut stream: TcpStream,
426 context: Arc<dyn AsyncEngineContext>,
427) -> Result<()> {
428 if context.is_killed() || context.is_stopped() {
432 tracing::debug!("stream context killed or stopped; skipping server FIN wait");
433 return Ok(());
434 }
435
436 let mut buf = [0u8; 1024];
439 let deadline = Instant::now() + Duration::from_secs(10);
440 loop {
441 let n = time::timeout_at(deadline, stream.read(&mut buf))
442 .await
443 .inspect_err(|_| {
444 tracing::debug!("server did not close socket within the deadline");
445 })?
446 .inspect_err(|e| {
447 tracing::debug!(err = ?e, "failed to read from stream");
448 })?;
449 if n == 0 {
450 break;
452 }
453 }
454
455 Ok(())
456}
457
458async fn handle_reader(
459 framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
460 context: Arc<dyn AsyncEngineContext>,
461 alive_tx: tokio::sync::oneshot::Sender<()>,
462 cancellation_counter: Option<IntCounter>,
463) -> FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec> {
464 let mut framed_reader = framed_reader;
465 let mut alive_tx = alive_tx;
466 let mut cancellation_seen = false;
468 loop {
469 tokio::select! {
470 msg = framed_reader.next() => {
471 match msg {
472 Some(Ok(two_part_msg)) => {
473 match two_part_msg.optional_parts() {
474 (Some(bytes), None) => {
475 let msg = match serde_json::from_slice::<ControlMessage>(bytes) {
476 Ok(msg) => msg,
477 Err(e) => {
478 tracing::warn!(
479 err = ?e,
480 "invalid control message, closing connection"
481 );
482 cancellation_seen = true;
483 context.kill();
484 break;
485 }
486 };
487
488 match msg {
495 ControlMessage::Stop => {
496 cancellation_seen = true;
497 context.stop();
498 }
499 ControlMessage::Kill => {
500 cancellation_seen = true;
501 context.kill();
502 }
503 ControlMessage::Sentinel => {
504 tracing::warn!(
505 "unexpected sentinel on client reader, closing connection"
506 );
507 cancellation_seen = true;
508 context.kill();
509 break;
510 }
511 }
512 }
513 _ => {
514 tracing::warn!(
515 "unexpected non-control message on client reader, closing connection"
516 );
517 cancellation_seen = true;
518 context.kill();
519 break;
520 }
521 }
522 }
523 Some(Err(e)) => {
524 tracing::warn!(err = ?e, "tcp stream read error, closing connection");
527 cancellation_seen = true;
528 context.kill();
529 break;
530 }
531 None => {
532 tracing::debug!("tcp stream closed by server");
533 cancellation_seen = true;
534 break;
535 }
536 }
537 }
538 _ = alive_tx.closed() => {
539 break;
540 }
541 }
542 }
543 if cancellation_seen && let Some(counter) = &cancellation_counter {
544 counter.inc();
545 }
546 framed_reader
547}
548
549async fn handle_writer(
550 mut framed_writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
551 mut bytes_rx: tokio::sync::mpsc::Receiver<TwoPartMessage>,
552 alive_rx: tokio::sync::oneshot::Receiver<()>,
553 context: Arc<dyn AsyncEngineContext>,
554) -> Result<FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>> {
555 let killed = context.killed();
558 let stopped = context.stopped();
559 tokio::pin!(killed, stopped);
560
561 let mut send_sentinel = true;
563
564 loop {
565 let msg = tokio::select! {
566 biased;
567
568 _ = &mut killed => {
569 tracing::trace!("context kill signal received; shutting down");
570 send_sentinel = false;
571 break;
572 }
573
574 _ = &mut stopped => {
575 tracing::trace!("context stop signal received; shutting down");
576 send_sentinel = false;
577 break;
578 }
579
580 msg = bytes_rx.recv() => {
581 match msg {
582 Some(msg) => msg,
583 None => {
584 tracing::trace!("response channel closed; shutting down");
585 break;
586 }
587 }
588 }
589 };
590
591 if let Err(e) = framed_writer.send(msg).await {
592 tracing::trace!(
593 "failed to send message to network; possible disconnect: {:?}",
594 e
595 );
596 send_sentinel = false;
597 break;
598 }
599 }
600
601 if send_sentinel {
603 let message = serde_json::to_vec(&ControlMessage::Sentinel)?;
604 let msg = TwoPartMessage::from_header(message.into());
605 framed_writer.send(msg).await?;
606 }
607
608 drop(alive_rx);
609 Ok(framed_writer)
610}
611
612#[cfg(test)]
613mod tests {
614 use super::*;
615 use crate::pipeline::context::Controller;
616 use crate::pipeline::network::tcp::test_utils::create_tcp_pair;
617 use bytes::Bytes;
618 use futures::StreamExt;
619 use std::sync::Arc;
620 use tokio::io::{AsyncReadExt, AsyncWriteExt};
621 use tokio::net::TcpStream;
622 use tokio::sync::{mpsc, oneshot};
623 use tokio_util::codec::FramedRead;
624
625 struct WriterHarness {
626 server: tokio::net::TcpStream,
627 framed_writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
628 bytes_tx: mpsc::Sender<TwoPartMessage>,
629 bytes_rx: mpsc::Receiver<TwoPartMessage>,
630 alive_tx: oneshot::Sender<()>,
631 alive_rx: oneshot::Receiver<()>,
632 controller: Arc<Controller>,
633 }
634
635 async fn writer_harness() -> WriterHarness {
637 let (client, server) = create_tcp_pair().await;
638 let (_, write_half) = tokio::io::split(client);
639 let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
640
641 let (bytes_tx, bytes_rx) = mpsc::channel(64);
642 let (alive_tx, alive_rx) = oneshot::channel::<()>();
643 let controller = Arc::new(Controller::default());
644
645 WriterHarness {
646 server,
647 framed_writer,
648 bytes_tx,
649 bytes_rx,
650 alive_tx,
651 alive_rx,
652 controller,
653 }
654 }
655
656 async fn recv_msg(reader: &mut FramedRead<TcpStream, TwoPartCodec>) -> TwoPartMessage {
657 reader
658 .next()
659 .await
660 .expect("expected message")
661 .expect("failed to decode message")
662 }
663
664 fn assert_data_only_message(msg: TwoPartMessage, expected: &[u8]) {
665 let (header, data) = msg.optional_parts();
666 assert!(header.is_none(), "data-only message should not have header");
667 assert_eq!(
668 data.expect("data payload missing").as_ref(),
669 expected,
670 "data payload should match"
671 );
672 }
673
674 fn assert_header_only_message(msg: TwoPartMessage, expected: &[u8]) {
675 let (header, data) = msg.optional_parts();
676 assert!(data.is_none(), "header-only message should not carry data");
677 assert_eq!(
678 header.expect("header missing").as_ref(),
679 expected,
680 "header payload should match"
681 );
682 }
683
684 fn assert_header_and_data_message(
685 msg: TwoPartMessage,
686 expected_header: &[u8],
687 expected_data: &[u8],
688 ) {
689 let (header, data) = msg.optional_parts();
690 assert_eq!(
691 header.expect("header missing").as_ref(),
692 expected_header,
693 "header payload should match"
694 );
695 assert_eq!(
696 data.expect("data missing").as_ref(),
697 expected_data,
698 "data payload should match"
699 );
700 }
701
702 fn assert_sentinel_message(msg: TwoPartMessage) {
703 let (header, data) = msg.optional_parts();
704 assert!(data.is_none(), "sentinel should not include a data section");
705 let expected_sentinel = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
706 assert_eq!(
707 header.expect("sentinel header missing").as_ref(),
708 expected_sentinel.as_slice(),
709 "sentinel header should match serialized ControlMessage::Sentinel"
710 );
711 }
712
713 #[tokio::test]
715 async fn test_handle_writer_forwards_messages() {
716 let WriterHarness {
717 server,
718 framed_writer,
719 bytes_tx,
720 bytes_rx,
721 alive_rx,
722 controller,
723 ..
724 } = writer_harness().await;
725
726 let test_msg = TwoPartMessage::from_data(Bytes::from("test data"));
728 bytes_tx.send(test_msg).await.unwrap();
729
730 drop(bytes_tx);
732
733 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
734
735 assert!(result.is_ok());
736
737 let mut reader = FramedRead::new(server, TwoPartCodec::default());
739
740 let msg = recv_msg(&mut reader).await;
741 assert_data_only_message(msg, b"test data");
742
743 let sentinel = recv_msg(&mut reader).await;
744 assert_sentinel_message(sentinel);
745 }
746
747 #[tokio::test]
749 async fn test_handle_writer_sends_sentinel_on_normal_closure() {
750 let WriterHarness {
751 mut server,
752 framed_writer,
753 bytes_tx,
754 bytes_rx,
755 alive_rx,
756 controller,
757 ..
758 } = writer_harness().await;
759
760 drop(bytes_tx);
762
763 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
764
765 assert!(result.is_ok());
766
767 let mut buffer = vec![0u8; 1024];
769 let n = server.read(&mut buffer).await.unwrap();
770
771 assert!(n > 0, "Expected sentinel to be written to the TCP stream");
773
774 let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
776 assert!(
777 buffer[..n]
778 .windows(sentinel_json.len())
779 .any(|w| w == sentinel_json.as_slice()),
780 "Buffer should contain sentinel message. Buffer: {:?}",
781 String::from_utf8_lossy(&buffer[..n])
782 );
783 }
784
785 #[tokio::test]
787 async fn test_handle_writer_no_sentinel_on_context_killed() {
788 let WriterHarness {
789 mut server,
790 framed_writer,
791 bytes_rx,
792 alive_rx,
793 controller,
794 ..
795 } = writer_harness().await;
796
797 controller.kill();
799
800 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
801
802 assert!(result.is_ok());
803
804 drop(result);
807
808 let mut buffer = vec![0u8; 1024];
810 let n = server.read(&mut buffer).await.unwrap();
811
812 let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
814 assert!(
815 n == 0
816 || !buffer[..n]
817 .windows(sentinel_json.len())
818 .any(|w| w == sentinel_json.as_slice()),
819 "Buffer should NOT contain sentinel message when context is killed"
820 );
821 }
822
823 #[tokio::test]
825 async fn test_handle_writer_no_sentinel_on_context_stopped() {
826 let WriterHarness {
827 mut server,
828 framed_writer,
829 bytes_rx,
830 alive_rx,
831 controller,
832 ..
833 } = writer_harness().await;
834
835 controller.stop();
837
838 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
839
840 assert!(result.is_ok());
841
842 drop(result);
845
846 let mut buffer = vec![0u8; 1024];
848 let n = server.read(&mut buffer).await.unwrap();
849
850 let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
852 assert!(
853 n == 0
854 || !buffer[..n]
855 .windows(sentinel_json.len())
856 .any(|w| w == sentinel_json.as_slice()),
857 "Buffer should NOT contain sentinel message when context is stopped"
858 );
859 }
860
861 #[tokio::test]
863 async fn test_handle_writer_multiple_messages() {
864 let WriterHarness {
865 server,
866 framed_writer,
867 bytes_tx,
868 bytes_rx,
869 alive_rx,
870 controller,
871 ..
872 } = writer_harness().await;
873
874 for i in 0..5 {
876 let test_msg = TwoPartMessage::from_data(Bytes::from(format!("message {}", i)));
877 bytes_tx.send(test_msg).await.unwrap();
878 }
879
880 drop(bytes_tx);
882
883 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
884
885 assert!(result.is_ok());
886
887 let mut reader = FramedRead::new(server, TwoPartCodec::default());
889 for i in 0..5 {
890 let msg = recv_msg(&mut reader).await;
891 assert_data_only_message(msg, format!("message {}", i).as_bytes());
892 }
893
894 let sentinel = recv_msg(&mut reader).await;
895 assert_sentinel_message(sentinel);
896 }
897
898 #[tokio::test]
900 async fn test_handle_writer_drops_alive_rx() {
901 let WriterHarness {
902 framed_writer,
903 bytes_tx,
904 bytes_rx,
905 alive_tx,
906 alive_rx,
907 controller,
908 ..
909 } = writer_harness().await;
910
911 drop(bytes_tx);
913
914 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
915
916 assert!(result.is_ok());
917
918 assert!(alive_tx.is_closed());
920 }
921
922 #[tokio::test]
924 async fn test_handle_writer_header_only_messages() {
925 let WriterHarness {
926 server,
927 framed_writer,
928 bytes_tx,
929 bytes_rx,
930 alive_rx,
931 controller,
932 ..
933 } = writer_harness().await;
934
935 let header_msg = TwoPartMessage::from_header(Bytes::from("header content"));
937 bytes_tx.send(header_msg).await.unwrap();
938
939 drop(bytes_tx);
941
942 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
943
944 assert!(result.is_ok());
945
946 let mut reader = FramedRead::new(server, TwoPartCodec::default());
947
948 let header_msg = recv_msg(&mut reader).await;
949 assert_header_only_message(header_msg, b"header content");
950
951 let sentinel = recv_msg(&mut reader).await;
952 assert_sentinel_message(sentinel);
953 }
954
955 #[tokio::test]
957 async fn test_handle_writer_mixed_messages() {
958 let WriterHarness {
959 server,
960 framed_writer,
961 bytes_tx,
962 bytes_rx,
963 alive_rx,
964 controller,
965 ..
966 } = writer_harness().await;
967
968 bytes_tx
970 .send(TwoPartMessage::from_header(Bytes::from("header1")))
971 .await
972 .unwrap();
973 bytes_tx
974 .send(TwoPartMessage::from_data(Bytes::from("data1")))
975 .await
976 .unwrap();
977 bytes_tx
978 .send(TwoPartMessage::from_parts(
979 Bytes::from("header2"),
980 Bytes::from("data2"),
981 ))
982 .await
983 .unwrap();
984
985 drop(bytes_tx);
987
988 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
989
990 assert!(result.is_ok());
991
992 let mut reader = FramedRead::new(server, TwoPartCodec::default());
993
994 let first = recv_msg(&mut reader).await;
995 assert_header_only_message(first, b"header1");
996
997 let second = recv_msg(&mut reader).await;
998 assert_data_only_message(second, b"data1");
999
1000 let third = recv_msg(&mut reader).await;
1001 assert_header_and_data_message(third, b"header2", b"data2");
1002
1003 let sentinel = recv_msg(&mut reader).await;
1004 assert_sentinel_message(sentinel);
1005 }
1006
1007 #[tokio::test]
1009 async fn test_wait_for_server_shutdown_skips_terminal_context() {
1010 for action in [Controller::kill as fn(&Controller), Controller::stop] {
1011 let (client, _server) = create_tcp_pair().await;
1012 let controller = Arc::new(Controller::default());
1013 action(&controller);
1014
1015 let context: Arc<dyn AsyncEngineContext> = controller;
1016 let result = tokio::time::timeout(
1017 std::time::Duration::from_millis(50),
1018 wait_for_server_shutdown(client, context),
1019 )
1020 .await;
1021
1022 assert!(result.is_ok(), "terminal context should not wait for FIN");
1023 assert!(
1024 result.unwrap().is_ok(),
1025 "terminal context shutdown should succeed"
1026 );
1027 }
1028 }
1029
1030 #[tokio::test]
1032 async fn test_connection_monitor_skips_fin_wait_after_read_error_kills_context() {
1033 let (client, mut server) = create_tcp_pair().await;
1034 let (read_half, write_half) = tokio::io::split(client);
1035 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1036 let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
1037 let (_bytes_tx, bytes_rx) = mpsc::channel(64);
1038 let (alive_tx, alive_rx) = oneshot::channel::<()>();
1039 let controller = Arc::new(Controller::default());
1040
1041 let reader_context = controller.clone();
1042 let reader_task = tokio::spawn(async move {
1043 handle_reader(framed_reader, reader_context, alive_tx, None).await
1044 });
1045 let writer_context = controller.clone();
1046 let writer_task = tokio::spawn(async move {
1047 handle_writer(framed_writer, bytes_rx, alive_rx, writer_context).await
1048 });
1049
1050 server.write_all(&[0xFF; 24]).await.unwrap();
1054
1055 let monitor_context: Arc<dyn AsyncEngineContext> = controller.clone();
1056 let result = tokio::time::timeout(
1057 std::time::Duration::from_millis(250),
1058 wait_for_connection_tasks(
1059 reader_task,
1060 writer_task,
1061 monitor_context,
1062 None,
1063 "test-subject".to_string(),
1064 ),
1065 )
1066 .await;
1067
1068 assert!(
1069 result.is_ok(),
1070 "connection monitor should not wait for the FIN deadline after read error"
1071 );
1072 assert!(result.unwrap().is_ok(), "connection monitor should succeed");
1073 assert!(
1074 controller.is_killed(),
1075 "read error should kill the stream context"
1076 );
1077 }
1078
1079 #[tokio::test]
1089 async fn test_connection_monitor_aborts_writer_when_reader_panics() {
1090 let reader_task: tokio::task::JoinHandle<
1094 FramedRead<ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
1095 > = tokio::spawn(async {
1096 panic!("simulated reader panic to trigger JoinError");
1097 });
1098
1099 let writer_task: tokio::task::JoinHandle<
1104 Result<FramedWrite<WriteHalf<tokio::net::TcpStream>, TwoPartCodec>>,
1105 > = tokio::spawn(async {
1106 std::future::pending::<()>().await;
1107 unreachable!()
1108 });
1109
1110 let controller = Arc::new(Controller::default());
1111 let context: Arc<dyn AsyncEngineContext> = controller.clone();
1112
1113 let result = tokio::time::timeout(
1116 std::time::Duration::from_millis(250),
1117 wait_for_connection_tasks(
1118 reader_task,
1119 writer_task,
1120 context,
1121 None,
1122 "test-reader-panic".to_string(),
1123 ),
1124 )
1125 .await;
1126
1127 assert!(
1130 result.is_ok(),
1131 "wait_for_connection_tasks must return after reader panic, \
1132 not hang waiting on the writer"
1133 );
1134
1135 assert!(
1137 result.unwrap().is_err(),
1138 "reader panic should propagate as Err from wait_for_connection_tasks"
1139 );
1140 }
1141
1142 struct ReaderHarness {
1145 framed_server: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
1146 framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
1147 alive_tx: oneshot::Sender<()>,
1148 alive_rx: oneshot::Receiver<()>,
1149 controller: Arc<Controller>,
1150 }
1151
1152 async fn reader_harness() -> ReaderHarness {
1154 let (client, server) = create_tcp_pair().await;
1155 let (read_half, _write_half) = tokio::io::split(client);
1156 let (_server_read, server_write) = tokio::io::split(server);
1157
1158 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1159 let framed_server = FramedWrite::new(server_write, TwoPartCodec::default());
1160 let (alive_tx, alive_rx) = oneshot::channel::<()>();
1161 let controller = Arc::new(Controller::default());
1162
1163 ReaderHarness {
1164 framed_server,
1165 framed_reader,
1166 alive_tx,
1167 alive_rx,
1168 controller,
1169 }
1170 }
1171
1172 fn control_message(msg: &ControlMessage) -> TwoPartMessage {
1173 let msg_bytes = serde_json::to_vec(msg).unwrap();
1174 TwoPartMessage::from_header(Bytes::from(msg_bytes))
1175 }
1176
1177 #[tokio::test]
1179 async fn test_handle_reader_stop_control_message() {
1180 let ReaderHarness {
1181 mut framed_server,
1182 framed_reader,
1183 alive_tx,
1184 alive_rx: _alive_rx,
1185 controller,
1186 } = reader_harness().await;
1187
1188 let controller_clone = controller.clone();
1190 let reader_handle = tokio::spawn(async move {
1191 handle_reader(framed_reader, controller_clone, alive_tx, None).await
1192 });
1193
1194 framed_server
1196 .send(control_message(&ControlMessage::Stop))
1197 .await
1198 .unwrap();
1199
1200 framed_server.close().await.unwrap();
1202
1203 let _ = reader_handle.await.unwrap();
1205
1206 assert!(
1208 controller.is_stopped(),
1209 "Controller should be stopped after receiving Stop message"
1210 );
1211 }
1212
1213 #[tokio::test]
1215 async fn test_handle_reader_kill_control_message() {
1216 let ReaderHarness {
1217 mut framed_server,
1218 framed_reader,
1219 alive_tx,
1220 alive_rx: _alive_rx,
1221 controller,
1222 } = reader_harness().await;
1223
1224 let controller_clone = controller.clone();
1226 let reader_handle = tokio::spawn(async move {
1227 handle_reader(framed_reader, controller_clone, alive_tx, None).await
1228 });
1229
1230 framed_server
1232 .send(control_message(&ControlMessage::Kill))
1233 .await
1234 .unwrap();
1235
1236 framed_server.close().await.unwrap();
1238
1239 let _ = reader_handle.await.unwrap();
1241
1242 assert!(
1244 controller.is_killed(),
1245 "Controller should be killed after receiving Kill message"
1246 );
1247 }
1248
1249 #[tokio::test]
1251 async fn test_handle_reader_exits_on_alive_channel_closed() {
1252 let ReaderHarness {
1253 framed_reader,
1254 alive_tx,
1255 alive_rx,
1256 controller,
1257 ..
1258 } = reader_harness().await;
1259
1260 let reader_handle =
1262 tokio::spawn(
1263 async move { handle_reader(framed_reader, controller, alive_tx, None).await },
1264 );
1265
1266 drop(alive_rx);
1268
1269 let result = reader_handle.await;
1271
1272 assert!(
1273 result.is_ok(),
1274 "handle_reader should exit when alive channel is closed"
1275 );
1276 }
1277
1278 #[tokio::test]
1280 async fn test_handle_reader_exits_on_stream_closed() {
1281 let ReaderHarness {
1282 mut framed_server,
1283 framed_reader,
1284 alive_tx,
1285 alive_rx: _alive_rx,
1286 controller,
1287 } = reader_harness().await;
1288
1289 let reader_handle =
1291 tokio::spawn(
1292 async move { handle_reader(framed_reader, controller, alive_tx, None).await },
1293 );
1294
1295 framed_server.close().await.unwrap();
1297
1298 let result = tokio::time::timeout(std::time::Duration::from_secs(1), reader_handle).await;
1300
1301 assert!(
1302 result.is_ok(),
1303 "handle_reader should exit when stream is closed"
1304 );
1305 }
1306
1307 #[tokio::test]
1309 async fn test_handle_reader_multiple_control_messages() {
1310 let ReaderHarness {
1311 mut framed_server,
1312 framed_reader,
1313 alive_tx,
1314 alive_rx: _alive_rx,
1315 controller,
1316 } = reader_harness().await;
1317
1318 let controller_clone = controller.clone();
1320 let reader_handle = tokio::spawn(async move {
1321 handle_reader(framed_reader, controller_clone, alive_tx, None).await
1322 });
1323
1324 framed_server
1326 .send(control_message(&ControlMessage::Stop))
1327 .await
1328 .unwrap();
1329 framed_server
1330 .send(control_message(&ControlMessage::Stop))
1331 .await
1332 .unwrap();
1333
1334 framed_server.close().await.unwrap();
1336
1337 let _ = reader_handle.await.unwrap();
1339
1340 assert!(
1342 controller.is_stopped(),
1343 "Controller should be stopped after receiving Stop messages"
1344 );
1345 }
1346
1347 #[tokio::test]
1349 async fn test_handle_reader_stop_then_kill() {
1350 let ReaderHarness {
1351 mut framed_server,
1352 framed_reader,
1353 alive_tx,
1354 alive_rx: _alive_rx,
1355 controller,
1356 } = reader_harness().await;
1357
1358 let controller_clone = controller.clone();
1360 let reader_handle = tokio::spawn(async move {
1361 handle_reader(framed_reader, controller_clone, alive_tx, None).await
1362 });
1363
1364 framed_server
1366 .send(control_message(&ControlMessage::Stop))
1367 .await
1368 .unwrap();
1369 framed_server
1370 .send(control_message(&ControlMessage::Kill))
1371 .await
1372 .unwrap();
1373
1374 framed_server.close().await.unwrap();
1376
1377 let _ = reader_handle.await.unwrap();
1379
1380 assert!(
1382 controller.is_killed(),
1383 "Controller should be killed after receiving Kill message"
1384 );
1385 }
1386
1387 #[tokio::test]
1389 async fn test_handle_reader_increments_cancellation_counter_on_read_error() {
1390 let ReaderHarness {
1391 framed_server,
1392 framed_reader,
1393 alive_tx,
1394 alive_rx: _alive_rx,
1395 controller,
1396 } = reader_harness().await;
1397 let cancellation_counter = IntCounter::new(
1398 "tcp_client_reader_read_error_cancellations_test",
1399 "test cancellation counter",
1400 )
1401 .unwrap();
1402
1403 let counter_clone = cancellation_counter.clone();
1404 let controller_clone = controller.clone();
1405 let reader_handle = tokio::spawn(async move {
1406 handle_reader(
1407 framed_reader,
1408 controller_clone,
1409 alive_tx,
1410 Some(counter_clone),
1411 )
1412 .await
1413 });
1414
1415 let mut raw_writer = framed_server.into_inner();
1416 raw_writer.write_all(&[0u8; 8]).await.unwrap();
1417 raw_writer.shutdown().await.unwrap();
1418
1419 let _ = reader_handle.await.unwrap();
1420
1421 assert!(
1422 controller.is_killed(),
1423 "Controller should be killed after TCP stream read error"
1424 );
1425 assert_eq!(
1426 cancellation_counter.get(),
1427 1,
1428 "read-error close should increment cancellation metric once"
1429 );
1430 }
1431
1432 async fn run_reader_with(
1435 msg: TwoPartMessage,
1436 counter_name: &str,
1437 ) -> (Arc<Controller>, IntCounter) {
1438 let ReaderHarness {
1439 mut framed_server,
1440 framed_reader,
1441 alive_tx,
1442 alive_rx: _alive_rx,
1443 controller,
1444 } = reader_harness().await;
1445 let counter = IntCounter::new(counter_name, "test counter").unwrap();
1446
1447 let counter_clone = counter.clone();
1448 let controller_clone = controller.clone();
1449 let reader_handle = tokio::spawn(async move {
1450 handle_reader(
1451 framed_reader,
1452 controller_clone,
1453 alive_tx,
1454 Some(counter_clone),
1455 )
1456 .await
1457 });
1458
1459 framed_server.send(msg).await.unwrap();
1460 let _ = reader_handle.await.unwrap();
1461
1462 (controller, counter)
1463 }
1464
1465 #[tokio::test]
1471 async fn test_handle_reader_kills_on_protocol_violations() {
1472 let cases: Vec<(&str, TwoPartMessage)> = vec![
1473 (
1474 "invalid control bytes",
1475 TwoPartMessage::from_header(Bytes::from_static(b"not a valid control message")),
1476 ),
1477 (
1478 "sentinel from server",
1479 control_message(&ControlMessage::Sentinel),
1480 ),
1481 (
1482 "non-control (data-only)",
1483 TwoPartMessage::from_data(Bytes::from_static(b"unexpected payload")),
1484 ),
1485 ];
1486
1487 for (i, (label, msg)) in cases.into_iter().enumerate() {
1488 let counter_name = format!("tcp_client_reader_protocol_violation_test_{i}");
1489 let (controller, counter) = run_reader_with(msg, &counter_name).await;
1490 assert!(
1491 controller.is_killed(),
1492 "{label}: should kill stream context"
1493 );
1494 assert_eq!(counter.get(), 1, "{label}: should be counted once");
1495 }
1496 }
1497
1498 struct RequestReaderHarness {
1501 framed_server: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
1502 framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
1503 bytes_tx: mpsc::Sender<Bytes>,
1504 bytes_rx: mpsc::Receiver<Bytes>,
1505 controller: Arc<Controller>,
1506 }
1507
1508 async fn request_reader_harness() -> RequestReaderHarness {
1509 let (client, server) = create_tcp_pair().await;
1510 let (read_half, _write_half) = tokio::io::split(client);
1511 let (_server_read, server_write) = tokio::io::split(server);
1512
1513 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1514 let framed_server = FramedWrite::new(server_write, TwoPartCodec::default());
1515 let (bytes_tx, bytes_rx) = mpsc::channel::<Bytes>(64);
1516 let controller = Arc::new(Controller::default());
1517
1518 RequestReaderHarness {
1519 framed_server,
1520 framed_reader,
1521 bytes_tx,
1522 bytes_rx,
1523 controller,
1524 }
1525 }
1526
1527 #[tokio::test]
1529 async fn test_handle_request_reader_stop_control_message() {
1530 let RequestReaderHarness {
1531 mut framed_server,
1532 framed_reader,
1533 bytes_tx,
1534 bytes_rx: _bytes_rx,
1535 controller,
1536 } = request_reader_harness().await;
1537
1538 let counter = IntCounter::new("tcp_request_reader_stop_test", "test counter").unwrap();
1539
1540 let counter_clone = counter.clone();
1541 let controller_clone = controller.clone();
1542 let handle = tokio::spawn(async move {
1543 handle_request_reader(
1544 framed_reader,
1545 bytes_tx,
1546 controller_clone,
1547 Some(counter_clone),
1548 )
1549 .await
1550 });
1551
1552 framed_server
1553 .send(control_message(&ControlMessage::Stop))
1554 .await
1555 .unwrap();
1556
1557 handle.await.unwrap();
1558
1559 assert!(controller.is_stopped(), "Stop should call context.stop()");
1560 assert!(!controller.is_killed(), "Stop should not kill the context");
1561 assert_eq!(counter.get(), 1, "cancellation counter should increment");
1562 }
1563
1564 #[tokio::test]
1566 async fn test_handle_request_reader_kill_control_message() {
1567 let RequestReaderHarness {
1568 mut framed_server,
1569 framed_reader,
1570 bytes_tx,
1571 bytes_rx: _bytes_rx,
1572 controller,
1573 } = request_reader_harness().await;
1574
1575 let counter = IntCounter::new("tcp_request_reader_kill_test", "test counter").unwrap();
1576
1577 let counter_clone = counter.clone();
1578 let controller_clone = controller.clone();
1579 let handle = tokio::spawn(async move {
1580 handle_request_reader(
1581 framed_reader,
1582 bytes_tx,
1583 controller_clone,
1584 Some(counter_clone),
1585 )
1586 .await
1587 });
1588
1589 framed_server
1590 .send(control_message(&ControlMessage::Kill))
1591 .await
1592 .unwrap();
1593
1594 handle.await.unwrap();
1595
1596 assert!(controller.is_killed(), "Kill should call context.kill()");
1597 assert_eq!(counter.get(), 1, "cancellation counter should increment");
1598 }
1599
1600 #[tokio::test]
1602 async fn test_handle_request_reader_sentinel_control_message() {
1603 let RequestReaderHarness {
1604 mut framed_server,
1605 framed_reader,
1606 bytes_tx,
1607 mut bytes_rx,
1608 controller,
1609 } = request_reader_harness().await;
1610
1611 let counter = IntCounter::new("tcp_request_reader_sentinel_test", "test counter").unwrap();
1612
1613 let counter_clone = counter.clone();
1614 let controller_clone = controller.clone();
1615 let handle = tokio::spawn(async move {
1616 handle_request_reader(
1617 framed_reader,
1618 bytes_tx,
1619 controller_clone,
1620 Some(counter_clone),
1621 )
1622 .await
1623 });
1624
1625 framed_server
1626 .send(control_message(&ControlMessage::Sentinel))
1627 .await
1628 .unwrap();
1629
1630 handle.await.unwrap();
1631
1632 assert!(
1633 !controller.is_stopped(),
1634 "Sentinel must not stop the context"
1635 );
1636 assert!(
1637 !controller.is_killed(),
1638 "Sentinel must not kill the context"
1639 );
1640 assert_eq!(counter.get(), 0, "Sentinel must not increment counter");
1641 assert!(
1642 bytes_rx.recv().await.is_none(),
1643 "bytes_tx should be dropped on exit"
1644 );
1645 }
1646
1647 #[tokio::test]
1650 async fn test_handle_request_reader_forwards_data() {
1651 let RequestReaderHarness {
1652 mut framed_server,
1653 framed_reader,
1654 bytes_tx,
1655 mut bytes_rx,
1656 controller,
1657 } = request_reader_harness().await;
1658
1659 let controller_clone = controller.clone();
1660 let handle = tokio::spawn(async move {
1661 handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
1662 });
1663
1664 framed_server
1665 .send(TwoPartMessage::from_data(Bytes::from_static(b"hello")))
1666 .await
1667 .unwrap();
1668 framed_server
1669 .send(TwoPartMessage::from_data(Bytes::from_static(b"world")))
1670 .await
1671 .unwrap();
1672
1673 assert_eq!(bytes_rx.recv().await.unwrap().as_ref(), b"hello");
1674 assert_eq!(bytes_rx.recv().await.unwrap().as_ref(), b"world");
1675
1676 framed_server
1677 .send(control_message(&ControlMessage::Sentinel))
1678 .await
1679 .unwrap();
1680
1681 handle.await.unwrap();
1682 assert!(
1683 bytes_rx.recv().await.is_none(),
1684 "channel should close after Sentinel"
1685 );
1686 }
1687
1688 #[tokio::test]
1690 async fn test_handle_request_reader_exits_on_context_killed() {
1691 let RequestReaderHarness {
1692 framed_server: _framed_server,
1693 framed_reader,
1694 bytes_tx,
1695 bytes_rx: _bytes_rx,
1696 controller,
1697 } = request_reader_harness().await;
1698
1699 let controller_clone = controller.clone();
1700 let handle = tokio::spawn(async move {
1701 handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
1702 });
1703
1704 controller.kill();
1705
1706 let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1707 assert!(
1708 result.is_ok(),
1709 "handler should exit promptly on context.kill()"
1710 );
1711 }
1712
1713 #[tokio::test]
1715 async fn test_handle_request_reader_exits_on_context_stopped() {
1716 let RequestReaderHarness {
1717 framed_server: _framed_server,
1718 framed_reader,
1719 bytes_tx,
1720 bytes_rx: _bytes_rx,
1721 controller,
1722 } = request_reader_harness().await;
1723
1724 let controller_clone = controller.clone();
1725 let handle = tokio::spawn(async move {
1726 handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
1727 });
1728
1729 controller.stop();
1730
1731 let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1732 assert!(
1733 result.is_ok(),
1734 "handler should exit promptly on context.stop()"
1735 );
1736 }
1737
1738 #[tokio::test]
1743 async fn test_handle_request_reader_exits_on_stream_closed() {
1744 let RequestReaderHarness {
1745 mut framed_server,
1746 framed_reader,
1747 bytes_tx,
1748 mut bytes_rx,
1749 controller,
1750 } = request_reader_harness().await;
1751
1752 let counter =
1753 IntCounter::new("tcp_request_reader_eof_truncation_test", "test counter").unwrap();
1754
1755 let counter_clone = counter.clone();
1756 let controller_clone = controller.clone();
1757 let handle = tokio::spawn(async move {
1758 handle_request_reader(
1759 framed_reader,
1760 bytes_tx,
1761 controller_clone,
1762 Some(counter_clone),
1763 )
1764 .await
1765 });
1766
1767 framed_server.close().await.unwrap();
1768
1769 let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1770 assert!(result.is_ok(), "handler should exit on EOF");
1771 assert!(
1772 controller.is_killed(),
1773 "EOF before sentinel should kill the context (truncated input)"
1774 );
1775 assert_eq!(
1776 counter.get(),
1777 1,
1778 "EOF before sentinel should count as a cancellation"
1779 );
1780 assert!(
1781 bytes_rx.recv().await.is_none(),
1782 "bytes_tx should be dropped"
1783 );
1784 }
1785
1786 #[tokio::test]
1791 async fn test_handle_request_reader_exits_when_receiver_dropped() {
1792 let RequestReaderHarness {
1793 framed_server,
1794 framed_reader,
1795 bytes_tx,
1796 bytes_rx,
1797 controller,
1798 } = request_reader_harness().await;
1799
1800 let _framed_server = framed_server;
1802
1803 let counter =
1804 IntCounter::new("tcp_request_reader_receiver_drop_test", "test counter").unwrap();
1805
1806 let counter_clone = counter.clone();
1807 let controller_clone = controller.clone();
1808 let handle = tokio::spawn(async move {
1809 handle_request_reader(
1810 framed_reader,
1811 bytes_tx,
1812 controller_clone,
1813 Some(counter_clone),
1814 )
1815 .await
1816 });
1817
1818 drop(bytes_rx);
1820
1821 let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1822 assert!(
1823 result.is_ok(),
1824 "handler should exit promptly when the receiver is dropped"
1825 );
1826 assert!(
1827 !controller.is_killed() && !controller.is_stopped(),
1828 "consumer drop is not a cancellation"
1829 );
1830 assert_eq!(
1831 counter.get(),
1832 0,
1833 "consumer drop must not count as cancellation"
1834 );
1835 }
1836}