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 mut cancellation_seen = false;
254 loop {
255 tokio::select! {
256 biased;
257
258 _ = context.killed() => {
259 tracing::trace!("context kill signal received on request stream; shutting down");
260 break;
261 }
262
263 _ = context.stopped() => {
264 tracing::trace!("context stop signal received on request stream; shutting down");
265 break;
266 }
267
268 _ = bytes_tx.closed() => {
273 tracing::debug!("downstream consumer dropped; exiting request-stream reader");
274 break;
275 }
276
277 msg = framed_reader.next() => {
278 match msg {
279 Some(Ok(two_part_msg)) => match two_part_msg.into_message_type() {
280 TwoPartMessageType::HeaderOnly(header) => {
281 let ctrl = match serde_json::from_slice::<ControlMessage>(&header) {
282 Ok(c) => c,
283 Err(e) => {
284 tracing::warn!(
285 err = ?e,
286 "invalid control message, closing connection"
287 );
288 cancellation_seen = true;
289 context.kill();
290 break;
291 }
292 };
293 match ctrl {
294 ControlMessage::Stop => {
295 cancellation_seen = true;
296 context.stop();
297 break;
298 }
299 ControlMessage::Kill => {
300 cancellation_seen = true;
301 context.kill();
302 break;
303 }
304 ControlMessage::Sentinel => {
305 tracing::trace!("upstream signaled end of request stream");
306 break;
307 }
308 }
309 }
310 TwoPartMessageType::DataOnly(data) => {
311 if bytes_tx.send(data).await.is_err() {
312 tracing::debug!("downstream consumer dropped; exiting request-stream reader");
313 break;
314 }
315 }
316 _ => {
317 tracing::warn!("fatal error - unexpected message shape on request stream");
318 cancellation_seen = true;
319 context.kill();
320 break;
321 }
322 }
323 Some(Err(e)) => {
324 tracing::warn!("fatal error - failed to decode message on request stream: {e:?}");
325 cancellation_seen = true;
326 context.kill();
327 break;
328 }
329 None => {
330 tracing::warn!("request stream closed by upstream before sentinel; treating as truncated");
335 cancellation_seen = true;
336 context.kill();
337 break;
338 }
339 }
340 }
341 }
342 }
343
344 if cancellation_seen && let Some(counter) = &cancellation_counter {
345 counter.inc();
346 }
347
348 drop(bytes_tx);
351}
352
353async fn wait_for_connection_tasks(
354 reader_task: tokio::task::JoinHandle<FramedRead<ReadHalf<TcpStream>, TwoPartCodec>>,
355 writer_task: tokio::task::JoinHandle<Result<FramedWrite<WriteHalf<TcpStream>, TwoPartCodec>>>,
356 context: Arc<dyn AsyncEngineContext>,
357 peer_port: Option<u16>,
358 subject: String,
359) -> Result<()> {
360 let reader = match reader_task.await {
363 Ok(reader) => reader,
364 Err(reader_err) => {
365 writer_task.abort();
366 let _ = writer_task.await;
367 tracing::error!(
368 subject = %subject,
369 peer_port = ?peer_port,
370 err = ?reader_err,
371 "reader task failed to join"
372 );
373 return Err(reader_err.into());
374 }
375 };
376
377 let writer = match writer_task.await {
378 Ok(writer) => writer,
379 Err(writer_err) => {
380 tracing::error!(
381 subject = %subject,
382 peer_port = ?peer_port,
383 err = ?writer_err,
384 "writer task failed to join"
385 );
386 return Err(writer_err.into());
387 }
388 };
389
390 let reader = reader.into_inner();
391 let writer = match writer {
392 Ok(writer) => writer.into_inner(),
393 Err(e) => {
394 tracing::error!(
395 subject = %subject,
396 peer_port = ?peer_port,
397 err = ?e,
398 "writer task returned error"
399 );
400 return Err(e);
401 }
402 };
403
404 let stream = reader.unsplit(writer);
405 wait_for_server_shutdown(stream, context).await
406}
407
408async fn wait_for_server_shutdown(
409 mut stream: TcpStream,
410 context: Arc<dyn AsyncEngineContext>,
411) -> Result<()> {
412 if context.is_killed() || context.is_stopped() {
416 tracing::debug!("stream context killed or stopped; skipping server FIN wait");
417 return Ok(());
418 }
419
420 let mut buf = [0u8; 1024];
423 let deadline = Instant::now() + Duration::from_secs(10);
424 loop {
425 let n = time::timeout_at(deadline, stream.read(&mut buf))
426 .await
427 .inspect_err(|_| {
428 tracing::debug!("server did not close socket within the deadline");
429 })?
430 .inspect_err(|e| {
431 tracing::debug!(err = ?e, "failed to read from stream");
432 })?;
433 if n == 0 {
434 break;
436 }
437 }
438
439 Ok(())
440}
441
442async fn handle_reader(
443 framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
444 context: Arc<dyn AsyncEngineContext>,
445 alive_tx: tokio::sync::oneshot::Sender<()>,
446 cancellation_counter: Option<IntCounter>,
447) -> FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec> {
448 let mut framed_reader = framed_reader;
449 let mut alive_tx = alive_tx;
450 let mut cancellation_seen = false;
452 loop {
453 tokio::select! {
454 msg = framed_reader.next() => {
455 match msg {
456 Some(Ok(two_part_msg)) => {
457 match two_part_msg.optional_parts() {
458 (Some(bytes), None) => {
459 let msg = match serde_json::from_slice::<ControlMessage>(bytes) {
460 Ok(msg) => msg,
461 Err(e) => {
462 tracing::warn!(
463 err = ?e,
464 "invalid control message, closing connection"
465 );
466 cancellation_seen = true;
467 context.kill();
468 break;
469 }
470 };
471
472 match msg {
479 ControlMessage::Stop => {
480 cancellation_seen = true;
481 context.stop();
482 }
483 ControlMessage::Kill => {
484 cancellation_seen = true;
485 context.kill();
486 }
487 ControlMessage::Sentinel => {
488 tracing::warn!(
489 "unexpected sentinel on client reader, closing connection"
490 );
491 cancellation_seen = true;
492 context.kill();
493 break;
494 }
495 }
496 }
497 _ => {
498 tracing::warn!(
499 "unexpected non-control message on client reader, closing connection"
500 );
501 cancellation_seen = true;
502 context.kill();
503 break;
504 }
505 }
506 }
507 Some(Err(e)) => {
508 tracing::warn!(err = ?e, "tcp stream read error, closing connection");
511 cancellation_seen = true;
512 context.kill();
513 break;
514 }
515 None => {
516 tracing::debug!("tcp stream closed by server");
517 cancellation_seen = true;
518 break;
519 }
520 }
521 }
522 _ = alive_tx.closed() => {
523 break;
524 }
525 }
526 }
527 if cancellation_seen && let Some(counter) = &cancellation_counter {
528 counter.inc();
529 }
530 framed_reader
531}
532
533async fn handle_writer(
534 mut framed_writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
535 mut bytes_rx: tokio::sync::mpsc::Receiver<TwoPartMessage>,
536 alive_rx: tokio::sync::oneshot::Receiver<()>,
537 context: Arc<dyn AsyncEngineContext>,
538) -> Result<FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>> {
539 let mut send_sentinel = true;
541
542 loop {
543 let msg = tokio::select! {
544 biased;
545
546 _ = context.killed() => {
547 tracing::trace!("context kill signal received; shutting down");
548 send_sentinel = false;
549 break;
550 }
551
552 _ = context.stopped() => {
553 tracing::trace!("context stop signal received; shutting down");
554 send_sentinel = false;
555 break;
556 }
557
558 msg = bytes_rx.recv() => {
559 match msg {
560 Some(msg) => msg,
561 None => {
562 tracing::trace!("response channel closed; shutting down");
563 break;
564 }
565 }
566 }
567 };
568
569 if let Err(e) = framed_writer.send(msg).await {
570 tracing::trace!(
571 "failed to send message to network; possible disconnect: {:?}",
572 e
573 );
574 send_sentinel = false;
575 break;
576 }
577 }
578
579 if send_sentinel {
581 let message = serde_json::to_vec(&ControlMessage::Sentinel)?;
582 let msg = TwoPartMessage::from_header(message.into());
583 framed_writer.send(msg).await?;
584 }
585
586 drop(alive_rx);
587 Ok(framed_writer)
588}
589
590#[cfg(test)]
591mod tests {
592 use super::*;
593 use crate::pipeline::context::Controller;
594 use crate::pipeline::network::tcp::test_utils::create_tcp_pair;
595 use bytes::Bytes;
596 use futures::StreamExt;
597 use std::sync::Arc;
598 use tokio::io::{AsyncReadExt, AsyncWriteExt};
599 use tokio::net::TcpStream;
600 use tokio::sync::{mpsc, oneshot};
601 use tokio_util::codec::FramedRead;
602
603 struct WriterHarness {
604 server: tokio::net::TcpStream,
605 framed_writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
606 bytes_tx: mpsc::Sender<TwoPartMessage>,
607 bytes_rx: mpsc::Receiver<TwoPartMessage>,
608 alive_tx: oneshot::Sender<()>,
609 alive_rx: oneshot::Receiver<()>,
610 controller: Arc<Controller>,
611 }
612
613 async fn writer_harness() -> WriterHarness {
615 let (client, server) = create_tcp_pair().await;
616 let (_, write_half) = tokio::io::split(client);
617 let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
618
619 let (bytes_tx, bytes_rx) = mpsc::channel(64);
620 let (alive_tx, alive_rx) = oneshot::channel::<()>();
621 let controller = Arc::new(Controller::default());
622
623 WriterHarness {
624 server,
625 framed_writer,
626 bytes_tx,
627 bytes_rx,
628 alive_tx,
629 alive_rx,
630 controller,
631 }
632 }
633
634 async fn recv_msg(reader: &mut FramedRead<TcpStream, TwoPartCodec>) -> TwoPartMessage {
635 reader
636 .next()
637 .await
638 .expect("expected message")
639 .expect("failed to decode message")
640 }
641
642 fn assert_data_only_message(msg: TwoPartMessage, expected: &[u8]) {
643 let (header, data) = msg.optional_parts();
644 assert!(header.is_none(), "data-only message should not have header");
645 assert_eq!(
646 data.expect("data payload missing").as_ref(),
647 expected,
648 "data payload should match"
649 );
650 }
651
652 fn assert_header_only_message(msg: TwoPartMessage, expected: &[u8]) {
653 let (header, data) = msg.optional_parts();
654 assert!(data.is_none(), "header-only message should not carry data");
655 assert_eq!(
656 header.expect("header missing").as_ref(),
657 expected,
658 "header payload should match"
659 );
660 }
661
662 fn assert_header_and_data_message(
663 msg: TwoPartMessage,
664 expected_header: &[u8],
665 expected_data: &[u8],
666 ) {
667 let (header, data) = msg.optional_parts();
668 assert_eq!(
669 header.expect("header missing").as_ref(),
670 expected_header,
671 "header payload should match"
672 );
673 assert_eq!(
674 data.expect("data missing").as_ref(),
675 expected_data,
676 "data payload should match"
677 );
678 }
679
680 fn assert_sentinel_message(msg: TwoPartMessage) {
681 let (header, data) = msg.optional_parts();
682 assert!(data.is_none(), "sentinel should not include a data section");
683 let expected_sentinel = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
684 assert_eq!(
685 header.expect("sentinel header missing").as_ref(),
686 expected_sentinel.as_slice(),
687 "sentinel header should match serialized ControlMessage::Sentinel"
688 );
689 }
690
691 #[tokio::test]
693 async fn test_handle_writer_forwards_messages() {
694 let WriterHarness {
695 server,
696 framed_writer,
697 bytes_tx,
698 bytes_rx,
699 alive_rx,
700 controller,
701 ..
702 } = writer_harness().await;
703
704 let test_msg = TwoPartMessage::from_data(Bytes::from("test data"));
706 bytes_tx.send(test_msg).await.unwrap();
707
708 drop(bytes_tx);
710
711 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
712
713 assert!(result.is_ok());
714
715 let mut reader = FramedRead::new(server, TwoPartCodec::default());
717
718 let msg = recv_msg(&mut reader).await;
719 assert_data_only_message(msg, b"test data");
720
721 let sentinel = recv_msg(&mut reader).await;
722 assert_sentinel_message(sentinel);
723 }
724
725 #[tokio::test]
727 async fn test_handle_writer_sends_sentinel_on_normal_closure() {
728 let WriterHarness {
729 mut server,
730 framed_writer,
731 bytes_tx,
732 bytes_rx,
733 alive_rx,
734 controller,
735 ..
736 } = writer_harness().await;
737
738 drop(bytes_tx);
740
741 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
742
743 assert!(result.is_ok());
744
745 let mut buffer = vec![0u8; 1024];
747 let n = server.read(&mut buffer).await.unwrap();
748
749 assert!(n > 0, "Expected sentinel to be written to the TCP stream");
751
752 let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
754 assert!(
755 buffer[..n]
756 .windows(sentinel_json.len())
757 .any(|w| w == sentinel_json.as_slice()),
758 "Buffer should contain sentinel message. Buffer: {:?}",
759 String::from_utf8_lossy(&buffer[..n])
760 );
761 }
762
763 #[tokio::test]
765 async fn test_handle_writer_no_sentinel_on_context_killed() {
766 let WriterHarness {
767 mut server,
768 framed_writer,
769 bytes_rx,
770 alive_rx,
771 controller,
772 ..
773 } = writer_harness().await;
774
775 controller.kill();
777
778 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
779
780 assert!(result.is_ok());
781
782 drop(result);
785
786 let mut buffer = vec![0u8; 1024];
788 let n = server.read(&mut buffer).await.unwrap();
789
790 let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
792 assert!(
793 n == 0
794 || !buffer[..n]
795 .windows(sentinel_json.len())
796 .any(|w| w == sentinel_json.as_slice()),
797 "Buffer should NOT contain sentinel message when context is killed"
798 );
799 }
800
801 #[tokio::test]
803 async fn test_handle_writer_no_sentinel_on_context_stopped() {
804 let WriterHarness {
805 mut server,
806 framed_writer,
807 bytes_rx,
808 alive_rx,
809 controller,
810 ..
811 } = writer_harness().await;
812
813 controller.stop();
815
816 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
817
818 assert!(result.is_ok());
819
820 drop(result);
823
824 let mut buffer = vec![0u8; 1024];
826 let n = server.read(&mut buffer).await.unwrap();
827
828 let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
830 assert!(
831 n == 0
832 || !buffer[..n]
833 .windows(sentinel_json.len())
834 .any(|w| w == sentinel_json.as_slice()),
835 "Buffer should NOT contain sentinel message when context is stopped"
836 );
837 }
838
839 #[tokio::test]
841 async fn test_handle_writer_multiple_messages() {
842 let WriterHarness {
843 server,
844 framed_writer,
845 bytes_tx,
846 bytes_rx,
847 alive_rx,
848 controller,
849 ..
850 } = writer_harness().await;
851
852 for i in 0..5 {
854 let test_msg = TwoPartMessage::from_data(Bytes::from(format!("message {}", i)));
855 bytes_tx.send(test_msg).await.unwrap();
856 }
857
858 drop(bytes_tx);
860
861 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
862
863 assert!(result.is_ok());
864
865 let mut reader = FramedRead::new(server, TwoPartCodec::default());
867 for i in 0..5 {
868 let msg = recv_msg(&mut reader).await;
869 assert_data_only_message(msg, format!("message {}", i).as_bytes());
870 }
871
872 let sentinel = recv_msg(&mut reader).await;
873 assert_sentinel_message(sentinel);
874 }
875
876 #[tokio::test]
878 async fn test_handle_writer_drops_alive_rx() {
879 let WriterHarness {
880 framed_writer,
881 bytes_tx,
882 bytes_rx,
883 alive_tx,
884 alive_rx,
885 controller,
886 ..
887 } = writer_harness().await;
888
889 drop(bytes_tx);
891
892 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
893
894 assert!(result.is_ok());
895
896 assert!(alive_tx.is_closed());
898 }
899
900 #[tokio::test]
902 async fn test_handle_writer_header_only_messages() {
903 let WriterHarness {
904 server,
905 framed_writer,
906 bytes_tx,
907 bytes_rx,
908 alive_rx,
909 controller,
910 ..
911 } = writer_harness().await;
912
913 let header_msg = TwoPartMessage::from_header(Bytes::from("header content"));
915 bytes_tx.send(header_msg).await.unwrap();
916
917 drop(bytes_tx);
919
920 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
921
922 assert!(result.is_ok());
923
924 let mut reader = FramedRead::new(server, TwoPartCodec::default());
925
926 let header_msg = recv_msg(&mut reader).await;
927 assert_header_only_message(header_msg, b"header content");
928
929 let sentinel = recv_msg(&mut reader).await;
930 assert_sentinel_message(sentinel);
931 }
932
933 #[tokio::test]
935 async fn test_handle_writer_mixed_messages() {
936 let WriterHarness {
937 server,
938 framed_writer,
939 bytes_tx,
940 bytes_rx,
941 alive_rx,
942 controller,
943 ..
944 } = writer_harness().await;
945
946 bytes_tx
948 .send(TwoPartMessage::from_header(Bytes::from("header1")))
949 .await
950 .unwrap();
951 bytes_tx
952 .send(TwoPartMessage::from_data(Bytes::from("data1")))
953 .await
954 .unwrap();
955 bytes_tx
956 .send(TwoPartMessage::from_parts(
957 Bytes::from("header2"),
958 Bytes::from("data2"),
959 ))
960 .await
961 .unwrap();
962
963 drop(bytes_tx);
965
966 let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
967
968 assert!(result.is_ok());
969
970 let mut reader = FramedRead::new(server, TwoPartCodec::default());
971
972 let first = recv_msg(&mut reader).await;
973 assert_header_only_message(first, b"header1");
974
975 let second = recv_msg(&mut reader).await;
976 assert_data_only_message(second, b"data1");
977
978 let third = recv_msg(&mut reader).await;
979 assert_header_and_data_message(third, b"header2", b"data2");
980
981 let sentinel = recv_msg(&mut reader).await;
982 assert_sentinel_message(sentinel);
983 }
984
985 #[tokio::test]
987 async fn test_wait_for_server_shutdown_skips_terminal_context() {
988 for action in [Controller::kill as fn(&Controller), Controller::stop] {
989 let (client, _server) = create_tcp_pair().await;
990 let controller = Arc::new(Controller::default());
991 action(&controller);
992
993 let context: Arc<dyn AsyncEngineContext> = controller;
994 let result = tokio::time::timeout(
995 std::time::Duration::from_millis(50),
996 wait_for_server_shutdown(client, context),
997 )
998 .await;
999
1000 assert!(result.is_ok(), "terminal context should not wait for FIN");
1001 assert!(
1002 result.unwrap().is_ok(),
1003 "terminal context shutdown should succeed"
1004 );
1005 }
1006 }
1007
1008 #[tokio::test]
1010 async fn test_connection_monitor_skips_fin_wait_after_read_error_kills_context() {
1011 let (client, mut server) = create_tcp_pair().await;
1012 let (read_half, write_half) = tokio::io::split(client);
1013 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1014 let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
1015 let (_bytes_tx, bytes_rx) = mpsc::channel(64);
1016 let (alive_tx, alive_rx) = oneshot::channel::<()>();
1017 let controller = Arc::new(Controller::default());
1018
1019 let reader_context = controller.clone();
1020 let reader_task = tokio::spawn(async move {
1021 handle_reader(framed_reader, reader_context, alive_tx, None).await
1022 });
1023 let writer_context = controller.clone();
1024 let writer_task = tokio::spawn(async move {
1025 handle_writer(framed_writer, bytes_rx, alive_rx, writer_context).await
1026 });
1027
1028 server.write_all(&[0xFF; 24]).await.unwrap();
1032
1033 let monitor_context: Arc<dyn AsyncEngineContext> = controller.clone();
1034 let result = tokio::time::timeout(
1035 std::time::Duration::from_millis(250),
1036 wait_for_connection_tasks(
1037 reader_task,
1038 writer_task,
1039 monitor_context,
1040 None,
1041 "test-subject".to_string(),
1042 ),
1043 )
1044 .await;
1045
1046 assert!(
1047 result.is_ok(),
1048 "connection monitor should not wait for the FIN deadline after read error"
1049 );
1050 assert!(result.unwrap().is_ok(), "connection monitor should succeed");
1051 assert!(
1052 controller.is_killed(),
1053 "read error should kill the stream context"
1054 );
1055 }
1056
1057 #[tokio::test]
1067 async fn test_connection_monitor_aborts_writer_when_reader_panics() {
1068 let reader_task: tokio::task::JoinHandle<
1072 FramedRead<ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
1073 > = tokio::spawn(async {
1074 panic!("simulated reader panic to trigger JoinError");
1075 });
1076
1077 let writer_task: tokio::task::JoinHandle<
1082 Result<FramedWrite<WriteHalf<tokio::net::TcpStream>, TwoPartCodec>>,
1083 > = tokio::spawn(async {
1084 std::future::pending::<()>().await;
1085 unreachable!()
1086 });
1087
1088 let controller = Arc::new(Controller::default());
1089 let context: Arc<dyn AsyncEngineContext> = controller.clone();
1090
1091 let result = tokio::time::timeout(
1094 std::time::Duration::from_millis(250),
1095 wait_for_connection_tasks(
1096 reader_task,
1097 writer_task,
1098 context,
1099 None,
1100 "test-reader-panic".to_string(),
1101 ),
1102 )
1103 .await;
1104
1105 assert!(
1108 result.is_ok(),
1109 "wait_for_connection_tasks must return after reader panic, \
1110 not hang waiting on the writer"
1111 );
1112
1113 assert!(
1115 result.unwrap().is_err(),
1116 "reader panic should propagate as Err from wait_for_connection_tasks"
1117 );
1118 }
1119
1120 struct ReaderHarness {
1123 framed_server: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
1124 framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
1125 alive_tx: oneshot::Sender<()>,
1126 alive_rx: oneshot::Receiver<()>,
1127 controller: Arc<Controller>,
1128 }
1129
1130 async fn reader_harness() -> ReaderHarness {
1132 let (client, server) = create_tcp_pair().await;
1133 let (read_half, _write_half) = tokio::io::split(client);
1134 let (_server_read, server_write) = tokio::io::split(server);
1135
1136 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1137 let framed_server = FramedWrite::new(server_write, TwoPartCodec::default());
1138 let (alive_tx, alive_rx) = oneshot::channel::<()>();
1139 let controller = Arc::new(Controller::default());
1140
1141 ReaderHarness {
1142 framed_server,
1143 framed_reader,
1144 alive_tx,
1145 alive_rx,
1146 controller,
1147 }
1148 }
1149
1150 fn control_message(msg: &ControlMessage) -> TwoPartMessage {
1151 let msg_bytes = serde_json::to_vec(msg).unwrap();
1152 TwoPartMessage::from_header(Bytes::from(msg_bytes))
1153 }
1154
1155 #[tokio::test]
1157 async fn test_handle_reader_stop_control_message() {
1158 let ReaderHarness {
1159 mut framed_server,
1160 framed_reader,
1161 alive_tx,
1162 alive_rx: _alive_rx,
1163 controller,
1164 } = reader_harness().await;
1165
1166 let controller_clone = controller.clone();
1168 let reader_handle = tokio::spawn(async move {
1169 handle_reader(framed_reader, controller_clone, alive_tx, None).await
1170 });
1171
1172 framed_server
1174 .send(control_message(&ControlMessage::Stop))
1175 .await
1176 .unwrap();
1177
1178 framed_server.close().await.unwrap();
1180
1181 let _ = reader_handle.await.unwrap();
1183
1184 assert!(
1186 controller.is_stopped(),
1187 "Controller should be stopped after receiving Stop message"
1188 );
1189 }
1190
1191 #[tokio::test]
1193 async fn test_handle_reader_kill_control_message() {
1194 let ReaderHarness {
1195 mut framed_server,
1196 framed_reader,
1197 alive_tx,
1198 alive_rx: _alive_rx,
1199 controller,
1200 } = reader_harness().await;
1201
1202 let controller_clone = controller.clone();
1204 let reader_handle = tokio::spawn(async move {
1205 handle_reader(framed_reader, controller_clone, alive_tx, None).await
1206 });
1207
1208 framed_server
1210 .send(control_message(&ControlMessage::Kill))
1211 .await
1212 .unwrap();
1213
1214 framed_server.close().await.unwrap();
1216
1217 let _ = reader_handle.await.unwrap();
1219
1220 assert!(
1222 controller.is_killed(),
1223 "Controller should be killed after receiving Kill message"
1224 );
1225 }
1226
1227 #[tokio::test]
1229 async fn test_handle_reader_exits_on_alive_channel_closed() {
1230 let ReaderHarness {
1231 framed_reader,
1232 alive_tx,
1233 alive_rx,
1234 controller,
1235 ..
1236 } = reader_harness().await;
1237
1238 let reader_handle =
1240 tokio::spawn(
1241 async move { handle_reader(framed_reader, controller, alive_tx, None).await },
1242 );
1243
1244 drop(alive_rx);
1246
1247 let result = reader_handle.await;
1249
1250 assert!(
1251 result.is_ok(),
1252 "handle_reader should exit when alive channel is closed"
1253 );
1254 }
1255
1256 #[tokio::test]
1258 async fn test_handle_reader_exits_on_stream_closed() {
1259 let ReaderHarness {
1260 mut framed_server,
1261 framed_reader,
1262 alive_tx,
1263 alive_rx: _alive_rx,
1264 controller,
1265 } = reader_harness().await;
1266
1267 let reader_handle =
1269 tokio::spawn(
1270 async move { handle_reader(framed_reader, controller, alive_tx, None).await },
1271 );
1272
1273 framed_server.close().await.unwrap();
1275
1276 let result = tokio::time::timeout(std::time::Duration::from_secs(1), reader_handle).await;
1278
1279 assert!(
1280 result.is_ok(),
1281 "handle_reader should exit when stream is closed"
1282 );
1283 }
1284
1285 #[tokio::test]
1287 async fn test_handle_reader_multiple_control_messages() {
1288 let ReaderHarness {
1289 mut framed_server,
1290 framed_reader,
1291 alive_tx,
1292 alive_rx: _alive_rx,
1293 controller,
1294 } = reader_harness().await;
1295
1296 let controller_clone = controller.clone();
1298 let reader_handle = tokio::spawn(async move {
1299 handle_reader(framed_reader, controller_clone, alive_tx, None).await
1300 });
1301
1302 framed_server
1304 .send(control_message(&ControlMessage::Stop))
1305 .await
1306 .unwrap();
1307 framed_server
1308 .send(control_message(&ControlMessage::Stop))
1309 .await
1310 .unwrap();
1311
1312 framed_server.close().await.unwrap();
1314
1315 let _ = reader_handle.await.unwrap();
1317
1318 assert!(
1320 controller.is_stopped(),
1321 "Controller should be stopped after receiving Stop messages"
1322 );
1323 }
1324
1325 #[tokio::test]
1327 async fn test_handle_reader_stop_then_kill() {
1328 let ReaderHarness {
1329 mut framed_server,
1330 framed_reader,
1331 alive_tx,
1332 alive_rx: _alive_rx,
1333 controller,
1334 } = reader_harness().await;
1335
1336 let controller_clone = controller.clone();
1338 let reader_handle = tokio::spawn(async move {
1339 handle_reader(framed_reader, controller_clone, alive_tx, None).await
1340 });
1341
1342 framed_server
1344 .send(control_message(&ControlMessage::Stop))
1345 .await
1346 .unwrap();
1347 framed_server
1348 .send(control_message(&ControlMessage::Kill))
1349 .await
1350 .unwrap();
1351
1352 framed_server.close().await.unwrap();
1354
1355 let _ = reader_handle.await.unwrap();
1357
1358 assert!(
1360 controller.is_killed(),
1361 "Controller should be killed after receiving Kill message"
1362 );
1363 }
1364
1365 #[tokio::test]
1367 async fn test_handle_reader_increments_cancellation_counter_on_read_error() {
1368 let ReaderHarness {
1369 framed_server,
1370 framed_reader,
1371 alive_tx,
1372 alive_rx: _alive_rx,
1373 controller,
1374 } = reader_harness().await;
1375 let cancellation_counter = IntCounter::new(
1376 "tcp_client_reader_read_error_cancellations_test",
1377 "test cancellation counter",
1378 )
1379 .unwrap();
1380
1381 let counter_clone = cancellation_counter.clone();
1382 let controller_clone = controller.clone();
1383 let reader_handle = tokio::spawn(async move {
1384 handle_reader(
1385 framed_reader,
1386 controller_clone,
1387 alive_tx,
1388 Some(counter_clone),
1389 )
1390 .await
1391 });
1392
1393 let mut raw_writer = framed_server.into_inner();
1394 raw_writer.write_all(&[0u8; 8]).await.unwrap();
1395 raw_writer.shutdown().await.unwrap();
1396
1397 let _ = reader_handle.await.unwrap();
1398
1399 assert!(
1400 controller.is_killed(),
1401 "Controller should be killed after TCP stream read error"
1402 );
1403 assert_eq!(
1404 cancellation_counter.get(),
1405 1,
1406 "read-error close should increment cancellation metric once"
1407 );
1408 }
1409
1410 async fn run_reader_with(
1413 msg: TwoPartMessage,
1414 counter_name: &str,
1415 ) -> (Arc<Controller>, IntCounter) {
1416 let ReaderHarness {
1417 mut framed_server,
1418 framed_reader,
1419 alive_tx,
1420 alive_rx: _alive_rx,
1421 controller,
1422 } = reader_harness().await;
1423 let counter = IntCounter::new(counter_name, "test counter").unwrap();
1424
1425 let counter_clone = counter.clone();
1426 let controller_clone = controller.clone();
1427 let reader_handle = tokio::spawn(async move {
1428 handle_reader(
1429 framed_reader,
1430 controller_clone,
1431 alive_tx,
1432 Some(counter_clone),
1433 )
1434 .await
1435 });
1436
1437 framed_server.send(msg).await.unwrap();
1438 let _ = reader_handle.await.unwrap();
1439
1440 (controller, counter)
1441 }
1442
1443 #[tokio::test]
1449 async fn test_handle_reader_kills_on_protocol_violations() {
1450 let cases: Vec<(&str, TwoPartMessage)> = vec![
1451 (
1452 "invalid control bytes",
1453 TwoPartMessage::from_header(Bytes::from_static(b"not a valid control message")),
1454 ),
1455 (
1456 "sentinel from server",
1457 control_message(&ControlMessage::Sentinel),
1458 ),
1459 (
1460 "non-control (data-only)",
1461 TwoPartMessage::from_data(Bytes::from_static(b"unexpected payload")),
1462 ),
1463 ];
1464
1465 for (i, (label, msg)) in cases.into_iter().enumerate() {
1466 let counter_name = format!("tcp_client_reader_protocol_violation_test_{i}");
1467 let (controller, counter) = run_reader_with(msg, &counter_name).await;
1468 assert!(
1469 controller.is_killed(),
1470 "{label}: should kill stream context"
1471 );
1472 assert_eq!(counter.get(), 1, "{label}: should be counted once");
1473 }
1474 }
1475
1476 struct RequestReaderHarness {
1479 framed_server: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
1480 framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
1481 bytes_tx: mpsc::Sender<Bytes>,
1482 bytes_rx: mpsc::Receiver<Bytes>,
1483 controller: Arc<Controller>,
1484 }
1485
1486 async fn request_reader_harness() -> RequestReaderHarness {
1487 let (client, server) = create_tcp_pair().await;
1488 let (read_half, _write_half) = tokio::io::split(client);
1489 let (_server_read, server_write) = tokio::io::split(server);
1490
1491 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1492 let framed_server = FramedWrite::new(server_write, TwoPartCodec::default());
1493 let (bytes_tx, bytes_rx) = mpsc::channel::<Bytes>(64);
1494 let controller = Arc::new(Controller::default());
1495
1496 RequestReaderHarness {
1497 framed_server,
1498 framed_reader,
1499 bytes_tx,
1500 bytes_rx,
1501 controller,
1502 }
1503 }
1504
1505 #[tokio::test]
1507 async fn test_handle_request_reader_stop_control_message() {
1508 let RequestReaderHarness {
1509 mut framed_server,
1510 framed_reader,
1511 bytes_tx,
1512 bytes_rx: _bytes_rx,
1513 controller,
1514 } = request_reader_harness().await;
1515
1516 let counter = IntCounter::new("tcp_request_reader_stop_test", "test counter").unwrap();
1517
1518 let counter_clone = counter.clone();
1519 let controller_clone = controller.clone();
1520 let handle = tokio::spawn(async move {
1521 handle_request_reader(
1522 framed_reader,
1523 bytes_tx,
1524 controller_clone,
1525 Some(counter_clone),
1526 )
1527 .await
1528 });
1529
1530 framed_server
1531 .send(control_message(&ControlMessage::Stop))
1532 .await
1533 .unwrap();
1534
1535 handle.await.unwrap();
1536
1537 assert!(controller.is_stopped(), "Stop should call context.stop()");
1538 assert!(!controller.is_killed(), "Stop should not kill the context");
1539 assert_eq!(counter.get(), 1, "cancellation counter should increment");
1540 }
1541
1542 #[tokio::test]
1544 async fn test_handle_request_reader_kill_control_message() {
1545 let RequestReaderHarness {
1546 mut framed_server,
1547 framed_reader,
1548 bytes_tx,
1549 bytes_rx: _bytes_rx,
1550 controller,
1551 } = request_reader_harness().await;
1552
1553 let counter = IntCounter::new("tcp_request_reader_kill_test", "test counter").unwrap();
1554
1555 let counter_clone = counter.clone();
1556 let controller_clone = controller.clone();
1557 let handle = tokio::spawn(async move {
1558 handle_request_reader(
1559 framed_reader,
1560 bytes_tx,
1561 controller_clone,
1562 Some(counter_clone),
1563 )
1564 .await
1565 });
1566
1567 framed_server
1568 .send(control_message(&ControlMessage::Kill))
1569 .await
1570 .unwrap();
1571
1572 handle.await.unwrap();
1573
1574 assert!(controller.is_killed(), "Kill should call context.kill()");
1575 assert_eq!(counter.get(), 1, "cancellation counter should increment");
1576 }
1577
1578 #[tokio::test]
1580 async fn test_handle_request_reader_sentinel_control_message() {
1581 let RequestReaderHarness {
1582 mut framed_server,
1583 framed_reader,
1584 bytes_tx,
1585 mut bytes_rx,
1586 controller,
1587 } = request_reader_harness().await;
1588
1589 let counter = IntCounter::new("tcp_request_reader_sentinel_test", "test counter").unwrap();
1590
1591 let counter_clone = counter.clone();
1592 let controller_clone = controller.clone();
1593 let handle = tokio::spawn(async move {
1594 handle_request_reader(
1595 framed_reader,
1596 bytes_tx,
1597 controller_clone,
1598 Some(counter_clone),
1599 )
1600 .await
1601 });
1602
1603 framed_server
1604 .send(control_message(&ControlMessage::Sentinel))
1605 .await
1606 .unwrap();
1607
1608 handle.await.unwrap();
1609
1610 assert!(
1611 !controller.is_stopped(),
1612 "Sentinel must not stop the context"
1613 );
1614 assert!(
1615 !controller.is_killed(),
1616 "Sentinel must not kill the context"
1617 );
1618 assert_eq!(counter.get(), 0, "Sentinel must not increment counter");
1619 assert!(
1620 bytes_rx.recv().await.is_none(),
1621 "bytes_tx should be dropped on exit"
1622 );
1623 }
1624
1625 #[tokio::test]
1628 async fn test_handle_request_reader_forwards_data() {
1629 let RequestReaderHarness {
1630 mut framed_server,
1631 framed_reader,
1632 bytes_tx,
1633 mut bytes_rx,
1634 controller,
1635 } = request_reader_harness().await;
1636
1637 let controller_clone = controller.clone();
1638 let handle = tokio::spawn(async move {
1639 handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
1640 });
1641
1642 framed_server
1643 .send(TwoPartMessage::from_data(Bytes::from_static(b"hello")))
1644 .await
1645 .unwrap();
1646 framed_server
1647 .send(TwoPartMessage::from_data(Bytes::from_static(b"world")))
1648 .await
1649 .unwrap();
1650
1651 assert_eq!(bytes_rx.recv().await.unwrap().as_ref(), b"hello");
1652 assert_eq!(bytes_rx.recv().await.unwrap().as_ref(), b"world");
1653
1654 framed_server
1655 .send(control_message(&ControlMessage::Sentinel))
1656 .await
1657 .unwrap();
1658
1659 handle.await.unwrap();
1660 assert!(
1661 bytes_rx.recv().await.is_none(),
1662 "channel should close after Sentinel"
1663 );
1664 }
1665
1666 #[tokio::test]
1668 async fn test_handle_request_reader_exits_on_context_killed() {
1669 let RequestReaderHarness {
1670 framed_server: _framed_server,
1671 framed_reader,
1672 bytes_tx,
1673 bytes_rx: _bytes_rx,
1674 controller,
1675 } = request_reader_harness().await;
1676
1677 let controller_clone = controller.clone();
1678 let handle = tokio::spawn(async move {
1679 handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
1680 });
1681
1682 controller.kill();
1683
1684 let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1685 assert!(
1686 result.is_ok(),
1687 "handler should exit promptly on context.kill()"
1688 );
1689 }
1690
1691 #[tokio::test]
1693 async fn test_handle_request_reader_exits_on_context_stopped() {
1694 let RequestReaderHarness {
1695 framed_server: _framed_server,
1696 framed_reader,
1697 bytes_tx,
1698 bytes_rx: _bytes_rx,
1699 controller,
1700 } = request_reader_harness().await;
1701
1702 let controller_clone = controller.clone();
1703 let handle = tokio::spawn(async move {
1704 handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
1705 });
1706
1707 controller.stop();
1708
1709 let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1710 assert!(
1711 result.is_ok(),
1712 "handler should exit promptly on context.stop()"
1713 );
1714 }
1715
1716 #[tokio::test]
1721 async fn test_handle_request_reader_exits_on_stream_closed() {
1722 let RequestReaderHarness {
1723 mut framed_server,
1724 framed_reader,
1725 bytes_tx,
1726 mut bytes_rx,
1727 controller,
1728 } = request_reader_harness().await;
1729
1730 let counter =
1731 IntCounter::new("tcp_request_reader_eof_truncation_test", "test counter").unwrap();
1732
1733 let counter_clone = counter.clone();
1734 let controller_clone = controller.clone();
1735 let handle = tokio::spawn(async move {
1736 handle_request_reader(
1737 framed_reader,
1738 bytes_tx,
1739 controller_clone,
1740 Some(counter_clone),
1741 )
1742 .await
1743 });
1744
1745 framed_server.close().await.unwrap();
1746
1747 let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1748 assert!(result.is_ok(), "handler should exit on EOF");
1749 assert!(
1750 controller.is_killed(),
1751 "EOF before sentinel should kill the context (truncated input)"
1752 );
1753 assert_eq!(
1754 counter.get(),
1755 1,
1756 "EOF before sentinel should count as a cancellation"
1757 );
1758 assert!(
1759 bytes_rx.recv().await.is_none(),
1760 "bytes_tx should be dropped"
1761 );
1762 }
1763
1764 #[tokio::test]
1769 async fn test_handle_request_reader_exits_when_receiver_dropped() {
1770 let RequestReaderHarness {
1771 framed_server,
1772 framed_reader,
1773 bytes_tx,
1774 bytes_rx,
1775 controller,
1776 } = request_reader_harness().await;
1777
1778 let _framed_server = framed_server;
1780
1781 let counter =
1782 IntCounter::new("tcp_request_reader_receiver_drop_test", "test counter").unwrap();
1783
1784 let counter_clone = counter.clone();
1785 let controller_clone = controller.clone();
1786 let handle = tokio::spawn(async move {
1787 handle_request_reader(
1788 framed_reader,
1789 bytes_tx,
1790 controller_clone,
1791 Some(counter_clone),
1792 )
1793 .await
1794 });
1795
1796 drop(bytes_rx);
1798
1799 let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1800 assert!(
1801 result.is_ok(),
1802 "handler should exit promptly when the receiver is dropped"
1803 );
1804 assert!(
1805 !controller.is_killed() && !controller.is_stopped(),
1806 "consumer drop is not a cancellation"
1807 );
1808 assert_eq!(
1809 counter.get(),
1810 0,
1811 "consumer drop must not count as cancellation"
1812 );
1813 }
1814}