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