Skip to main content

dynamo_runtime/pipeline/network/tcp/
client.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use 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}; // Import SinkExt to use the `send` method
25
26#[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        // try to connect to the address; retry with linear backoff if AddrNotAvailable
46        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        // this is a oneshot channel that will be used to signal when the stream is closed
97        // when the stream sender is dropped, the bytes_rx will be closed and the forwarder task will exit
98        // the forwarder task will capture the alive_rx half of the oneshot channel; this will close the alive channel
99        // so the holder of the alive_tx half will be notified that the stream is closed; the alive_tx channel will be
100        // captured by the monitor task
101        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        // transport specific handshake message
111        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        // issue the the first tcp handshake message
127        framed_writer
128            .send(msg)
129            .await
130            .map_err(|e| error!("failed to send handshake: {:?}", e))?;
131
132        // set up the channel to send bytes to the transport layer
133        let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel(64);
134
135        // forwards the bytes send from this stream to the transport layer; hold the alive_rx half of the oneshot channel
136        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        // Spawn the connection monitor; errors are already logged inside
147        // wait_for_connection_tasks, so the Result is intentionally dropped.
148        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        // set up the prologue for the stream
160        // this might have transport specific metadata in the future
161        let prologue = Some(ResponseStreamPrologue { error: None });
162
163        // create the stream sender
164        let stream_sender = StreamSender {
165            tx: bytes_tx,
166            prologue,
167        };
168
169        Ok(stream_sender)
170    }
171
172    /// Symmetric to [`Self::create_response_stream`] for the request-stream half:
173    /// dial the upstream TCP server with `StreamType::Request`, then return a
174    /// [`StreamReceiver`] that yields the data frames the upstream pushes down.
175    ///
176    /// The request stream is unidirectional after the handshake: the write half
177    /// is dropped as soon as the `CallHomeHandshake` is sent, so the downstream
178    /// never writes anything back (no `Sentinel` ack). The spawned reader task
179    /// forwards `TwoPartMessage::DataOnly` payloads into the channel and
180    /// translates `ControlMessage::Stop` / `Kill` into context cancellation;
181    /// `ControlMessage::Sentinel` terminates the task cleanly. A TCP close
182    /// before any `Sentinel` is treated as a truncated input (cancellation +
183    /// `context.kill()`), and dropping the returned `StreamReceiver` also stops
184    /// the task.
185    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        // Request stream is unidirectional after the handshake: the downstream
230        // never writes again, so close the write half immediately.
231        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    // Only mark cancellation on fatal errors or explicit upstream cancellation.
253    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            // Downstream consumer dropped the StreamReceiver. Exit promptly
269            // instead of staying parked on `framed_reader.next()` until the
270            // socket closes — the data has nowhere to go. This is the consumer's
271            // own choice, so it is not a cancellation (no kill, no count).
272            _ = 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                        // Socket closed before a Sentinel/Stop/Kill: the request
331                        // input is truncated. Kill the context so the consumer
332                        // sees an aborted stream rather than a clean end, and
333                        // count it as a cancellation.
334                        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    // Dropping bytes_tx closes the receiver side, signaling end-of-stream to the
349    // engine consumer.
350    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    // Await the reader first and abort the writer on reader Err — the
361    // writer parks on `bytes_rx.recv()` and won't wake on its own.
362    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    // `handle_writer` skips the closing sentinel on both `killed` and
413    // `stopped`, so the server has nothing to react to in either case;
414    // sitting in the read loop until the 10 s deadline would be dead time.
415    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    // Await the tcp server to shutdown the socket connection, bounded by a
421    // timeout so normal sentinel shutdown cannot hang indefinitely.
422    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            // Server has closed (FIN)
435            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    // Set on every cancellation arm; counted once after the loop.
451    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                                // Stop/Kill intentionally do not `break`: the
473                                // reader keeps running so a later Kill can
474                                // upgrade an earlier Stop (and vice versa).
475                                // The loop still exits promptly via the
476                                // `alive_tx.closed()` arm once `handle_writer`
477                                // reacts to `context.stop()` / `context.kill()`.
478                                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                        // Kill the engine context so the producer stops
509                        // generating responses that can no longer be delivered.
510                        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    // Only send sentinel for normal channel closure
540    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    // Send sentinel only on normal closure
580    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    /// Creates a reusable writer harness with paired TCP streams and test channels.
614    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    /// Test that handle_writer forwards messages from the channel to the framed writer
692    #[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        // Send test messages
705        let test_msg = TwoPartMessage::from_data(Bytes::from("test data"));
706        bytes_tx.send(test_msg).await.unwrap();
707
708        // Close the sender to trigger normal termination
709        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        // Decode from server side to verify data and sentinel were sent
716        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    /// Test that handle_writer sends sentinel on normal channel closure
726    #[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        // Close the sender immediately to trigger normal termination
739        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        // Read from server side to verify sentinel was sent
746        let mut buffer = vec![0u8; 1024];
747        let n = server.read(&mut buffer).await.unwrap();
748
749        // Buffer should contain the sentinel message
750        assert!(n > 0, "Expected sentinel to be written to the TCP stream");
751
752        // Verify it contains the sentinel message by checking for the JSON
753        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    /// Test that handle_writer does NOT send sentinel when context is killed
764    #[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        // Kill the context
776        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 the writer to close the connection, then try to read. Otherwise,
783        // the test will hang on `server.read()`
784        drop(result);
785
786        // Read from server side - should get no sentinel
787        let mut buffer = vec![0u8; 1024];
788        let n = server.read(&mut buffer).await.unwrap();
789
790        // Buffer should be empty (no sentinel sent)
791        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    /// Test that handle_writer does NOT send sentinel when context is stopped
802    #[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        // Stop the context
814        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 the writer to close the connection, then try to read. Otherwise,
821        // the test will hang on `server.read()`
822        drop(result);
823
824        // Read from server side - should get no sentinel
825        let mut buffer = vec![0u8; 1024];
826        let n = server.read(&mut buffer).await.unwrap();
827
828        // Buffer should be empty (no sentinel sent)
829        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    /// Test that handle_writer handles multiple messages correctly
840    #[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        // Send multiple messages
853        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        // Close the sender to trigger normal termination
859        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        // Decode from server side to verify all messages plus sentinel
866        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    /// Test that alive_rx is dropped after handle_writer completes
877    #[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        // Close the sender to trigger normal termination
890        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        // alive_tx should now be closed because alive_rx was dropped
897        assert!(alive_tx.is_closed());
898    }
899
900    /// Test handle_writer with header-only messages (control messages)
901    #[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        // Send a header-only message
914        let header_msg = TwoPartMessage::from_header(Bytes::from("header content"));
915        bytes_tx.send(header_msg).await.unwrap();
916
917        // Close the sender
918        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    /// Test handle_writer with mixed header and data messages
934    #[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        // Send mixed messages
947        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        // Close the sender
964        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    /// Killed or stopped contexts skip the server FIN deadline.
986    #[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    /// Read error in the connection monitor kills the context and skips the FIN wait.
1009    #[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        // Bypass the codec and write a complete but invalid TwoPartCodec
1029        // header. This drives the client reader into Some(Err(_)) without
1030        // closing the server side of the socket.
1031        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    /// Reader-side panic must abort the writer and return promptly rather than
1058    /// hanging on `tokio::join!`. Locks in the fix added with this function's
1059    /// sequential-await + writer-abort behavior.
1060    ///
1061    /// Setup: spawn a reader task that panics immediately (so
1062    /// `reader_task.await` yields `Err(JoinError::panic)`), and a writer task
1063    /// that parks indefinitely waiting for application bytes (so without the
1064    /// abort, `tokio::join!` on the previous implementation would never wake).
1065    /// Expect: `wait_for_connection_tasks` returns Err within the timeout.
1066    #[tokio::test]
1067    async fn test_connection_monitor_aborts_writer_when_reader_panics() {
1068        // Reader task that panics immediately. The explicit JoinHandle type
1069        // pins the inferred return type to the one wait_for_connection_tasks
1070        // expects; `panic!` is type `!`, which coerces to that type.
1071        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        // Writer task that would block indefinitely waiting on application
1078        // bytes. Under the pre-fix `tokio::join!` implementation, this would
1079        // prevent the function from returning when the reader panicked.
1080        // After the fix, the abort drives this task to completion promptly.
1081        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        // 250 ms is generous — the abort + JoinHandle resolution should fire
1092        // sub-millisecond. We are checking for "doesn't hang", not "fast".
1093        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        // Outer timeout must not fire: the abort path must surface the reader
1106        // JoinError before the writer would have produced any bytes.
1107        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        // The inner result must be Err — the reader's JoinError propagates.
1114        assert!(
1115            result.unwrap().is_err(),
1116            "reader panic should propagate as Err from wait_for_connection_tasks"
1117        );
1118    }
1119
1120    // ==================== handle_reader tests ====================
1121
1122    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    /// Creates a reusable reader harness with paired TCP streams and test channels.
1131    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    /// Test that handle_reader handles Stop control message by calling context.stop()
1156    #[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        // Spawn the reader task
1167        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        // Send Stop control message from server
1173        framed_server
1174            .send(control_message(&ControlMessage::Stop))
1175            .await
1176            .unwrap();
1177
1178        // Close the framed server to signal EOF to the client
1179        framed_server.close().await.unwrap();
1180
1181        // Wait for reader to finish
1182        let _ = reader_handle.await.unwrap();
1183
1184        // Verify that stop was called on the controller
1185        assert!(
1186            controller.is_stopped(),
1187            "Controller should be stopped after receiving Stop message"
1188        );
1189    }
1190
1191    /// Test that handle_reader handles Kill control message by calling context.kill()
1192    #[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        // Spawn the reader task
1203        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        // Send Kill control message from server
1209        framed_server
1210            .send(control_message(&ControlMessage::Kill))
1211            .await
1212            .unwrap();
1213
1214        // Close the framed server to signal EOF to the client
1215        framed_server.close().await.unwrap();
1216
1217        // Wait for reader to finish
1218        let _ = reader_handle.await.unwrap();
1219
1220        // Verify that kill was called on the controller
1221        assert!(
1222            controller.is_killed(),
1223            "Controller should be killed after receiving Kill message"
1224        );
1225    }
1226
1227    /// Test that handle_reader exits when alive channel is closed
1228    #[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        // Spawn the reader task
1239        let reader_handle =
1240            tokio::spawn(
1241                async move { handle_reader(framed_reader, controller, alive_tx, None).await },
1242            );
1243
1244        // Drop the alive_rx to close the channel (simulating writer finishing)
1245        drop(alive_rx);
1246
1247        // Reader should exit due to alive channel closure
1248        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    /// Test that handle_reader exits when TCP stream is closed
1257    #[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        // Spawn the reader task
1268        let reader_handle =
1269            tokio::spawn(
1270                async move { handle_reader(framed_reader, controller, alive_tx, None).await },
1271            );
1272
1273        // Close the framed server to signal EOF to the client
1274        framed_server.close().await.unwrap();
1275
1276        // Reader should exit due to stream closure
1277        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    /// Test that handle_reader handles multiple control messages in sequence
1286    #[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        // Spawn the reader task
1297        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        // Send multiple Stop messages (first one will stop, subsequent ones are no-ops)
1303        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        // Close the framed server to signal EOF to the client
1313        framed_server.close().await.unwrap();
1314
1315        // Wait for reader to finish
1316        let _ = reader_handle.await.unwrap();
1317
1318        // Verify that stop was called
1319        assert!(
1320            controller.is_stopped(),
1321            "Controller should be stopped after receiving Stop messages"
1322        );
1323    }
1324
1325    /// Test handle_reader with Stop followed by Kill
1326    #[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        // Spawn the reader task
1337        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        // Send Stop first, then Kill
1343        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        // Close the framed server to signal EOF to the client
1353        framed_server.close().await.unwrap();
1354
1355        // Wait for reader to finish
1356        let _ = reader_handle.await.unwrap();
1357
1358        // Verify that kill was called (which sets killed state)
1359        assert!(
1360            controller.is_killed(),
1361            "Controller should be killed after receiving Kill message"
1362        );
1363    }
1364
1365    /// Read errors kill the context and are counted as cancellations.
1366    #[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    /// Drives `handle_reader` against a single message and returns the
1411    /// controller + cancellation counter for assertions.
1412    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    /// Each protocol-violating message variant must kill only this stream
1444    /// (controller killed, cancellation counted once) and never panic the
1445    /// worker. Covers the three non-read-error panic arms in `handle_reader`:
1446    /// undecodable control bytes, server-sent Sentinel, and non-control
1447    /// (data-only) messages.
1448    #[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    // ==================== handle_request_reader tests ====================
1477
1478    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    /// Receiving Stop calls context.stop(), increments the counter, and exits.
1506    #[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    /// Receiving Kill calls context.kill(), increments the counter, and exits.
1543    #[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    /// Receiving Sentinel exits cleanly without touching the context or counter.
1579    #[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    /// DataOnly frames are forwarded to bytes_tx; the loop continues until a
1626    /// terminator arrives (here, Sentinel).
1627    #[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    /// External context.kill() exits the reader without touching the wire.
1667    #[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    /// External context.stop() exits the reader without touching the wire.
1692    #[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    /// Socket EOF exits the reader and drops bytes_tx.
1717    /// EOF before a closing Sentinel is a truncated request input: the handler
1718    /// kills the context and counts a cancellation so the consumer sees an
1719    /// aborted stream rather than a clean end.
1720    #[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    /// Dropping the returned StreamReceiver makes the reader exit promptly via
1765    /// the `bytes_tx.closed()` arm, even while parked on the socket with no
1766    /// incoming frame. This is the consumer's own choice, so it is not counted
1767    /// as a cancellation and the context is left untouched.
1768    #[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        // Keep the socket open so the only exit path is the receiver drop.
1779        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 the consumer; the reader is parked on `framed_reader.next()`.
1797        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}