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 rustls::pki_types::ServerName;
8use tokio::io::AsyncReadExt;
9use tokio::{
10    net::TcpStream,
11    time::{self, Duration, Instant},
12};
13use tokio_rustls::TlsConnector;
14use tokio_util::codec::{FramedRead, FramedWrite};
15
16type BoxRead = Box<dyn tokio::io::AsyncRead + Unpin + Send>;
17type BoxWrite = Box<dyn tokio::io::AsyncWrite + Unpin + Send>;
18
19use prometheus::IntCounter;
20use tracing::Instrument;
21
22use super::{CallHomeHandshake, ControlMessage, TcpStreamConnectionInfo};
23use crate::engine::AsyncEngineContext;
24use crate::pipeline::network::{
25    ConnectionInfo, ResponseStreamPrologue, StreamReceiver, StreamSender,
26    codec::{TwoPartCodec, TwoPartMessage, TwoPartMessageType},
27    tcp::StreamType,
28};
29use anyhow::{Context, Result, anyhow as error}; // Import SinkExt to use the `send` method
30
31#[allow(dead_code)]
32pub struct TcpClient {
33    worker_id: String,
34}
35
36fn record_upstream_cancellation(context: &dyn AsyncEngineContext, signal: &'static str) {
37    tracing::info!(
38        target: "request_span",
39        {
40            request_id = context.id(),
41            { "cancellation.signal" } = signal,
42            { "cancellation.source" } = "upstream"
43        },
44        "request cancellation received"
45    );
46}
47
48impl Default for TcpClient {
49    fn default() -> Self {
50        TcpClient {
51            worker_id: uuid::Uuid::new_v4().to_string(),
52        }
53    }
54}
55
56impl TcpClient {
57    pub fn new(worker_id: String) -> Self {
58        TcpClient { worker_id }
59    }
60
61    async fn connect(address: &str) -> std::io::Result<TcpStream> {
62        // try to connect to the address; retry with linear backoff if AddrNotAvailable
63        let backoff = std::time::Duration::from_millis(200);
64        loop {
65            match TcpStream::connect(address).await {
66                Ok(socket) => {
67                    socket.set_nodelay(true)?;
68                    return Ok(socket);
69                }
70                Err(e) => {
71                    if e.kind() == std::io::ErrorKind::AddrNotAvailable {
72                        tracing::warn!("retry warning: failed to connect: {:?}", e);
73                        tokio::time::sleep(backoff).await;
74                    } else {
75                        return Err(e);
76                    }
77                }
78            }
79        }
80    }
81
82    /// Connect to `address` and return split, boxed read/write halves.
83    /// When any TCP TLS env var is configured the connection is upgraded to TLS,
84    /// using the process-global connector (built once from the environment).
85    async fn connect_and_split(address: &str) -> anyhow::Result<(BoxRead, BoxWrite)> {
86        Self::connect_and_split_with_connector(address, get_tls_connector()?.as_ref()).await
87    }
88
89    /// Like [`connect_and_split`], but with an explicitly supplied TLS connector
90    /// instead of the process-global `OnceCell`. Lets tests drive the real client
91    /// handshake + boxed-I/O path deterministically (the `OnceCell` can only be
92    /// initialized once per process). Mirrors the request-plane client's
93    /// `connect_with_connector`.
94    async fn connect_and_split_with_connector(
95        address: &str,
96        connector: Option<&TlsConnector>,
97    ) -> anyhow::Result<(BoxRead, BoxWrite)> {
98        let stream = TcpClient::connect(address).await?;
99        if let Some(connector) = connector {
100            let server_name = tls_server_name(address)?;
101            let tls_stream = tokio::time::timeout(
102                crate::tls_utils::handshake_timeout(),
103                connector.connect(server_name, stream),
104            )
105            .await
106            .with_context(|| format!("TLS handshake timed out connecting to {address}"))?
107            .with_context(|| format!("TLS handshake failed connecting to {address}"))?;
108            let (r, w) = tokio::io::split(tls_stream);
109            Ok((Box::new(r), Box::new(w)))
110        } else {
111            let (r, w) = tokio::io::split(stream);
112            Ok((Box::new(r), Box::new(w)))
113        }
114    }
115
116    pub async fn create_response_stream(
117        context: Arc<dyn AsyncEngineContext>,
118        info: ConnectionInfo,
119        cancellation_counter: Option<IntCounter>,
120    ) -> Result<StreamSender> {
121        let info =
122            TcpStreamConnectionInfo::try_from(info).context("tcp-stream-connection-info-error")?;
123        tracing::trace!("Creating response stream for {:?}", info);
124
125        if info.stream_type != StreamType::Response {
126            return Err(error!(
127                "Invalid stream type; TcpClient requires the stream type to be `response`; however {:?} was passed",
128                info.stream_type
129            ));
130        }
131
132        if info.context != context.id() {
133            return Err(error!(
134                "Invalid context; TcpClient requires the context to be {:?}; however {:?} was passed",
135                context.id(),
136                info.context
137            ));
138        }
139
140        let (read_half, write_half) = TcpClient::connect_and_split(&info.address).await?;
141
142        let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
143        let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
144
145        // this is a oneshot channel that will be used to signal when the stream is closed
146        // when the stream sender is dropped, the bytes_rx will be closed and the forwarder task will exit
147        // the forwarder task will capture the alive_rx half of the oneshot channel; this will close the alive channel
148        // so the holder of the alive_tx half will be notified that the stream is closed; the alive_tx channel will be
149        // captured by the monitor task
150        let (alive_tx, alive_rx) = tokio::sync::oneshot::channel::<()>();
151
152        let reader_span = tracing::Span::current();
153        let reader_task = tokio::spawn(
154            handle_reader(
155                framed_reader,
156                context.clone(),
157                alive_tx,
158                cancellation_counter,
159            )
160            .instrument(reader_span),
161        );
162
163        // transport specific handshake message
164        let handshake = CallHomeHandshake {
165            subject: info.subject.clone(),
166            stream_type: StreamType::Response,
167        };
168
169        let handshake_bytes = match serde_json::to_vec(&handshake) {
170            Ok(hb) => hb,
171            Err(err) => {
172                return Err(error!(
173                    "create_response_stream: Error converting CallHomeHandshake to JSON array: {err:#}"
174                ));
175            }
176        };
177        let msg = TwoPartMessage::from_header(handshake_bytes.into());
178
179        // issue the the first tcp handshake message
180        framed_writer
181            .send(msg)
182            .await
183            .map_err(|e| error!("failed to send handshake: {:?}", e))?;
184
185        // set up the channel to send bytes to the transport layer
186        let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel(64);
187
188        // forwards the bytes send from this stream to the transport layer; hold the alive_rx half of the oneshot channel
189        let writer_context = context.clone();
190        let writer_task = tokio::spawn(handle_writer(
191            framed_writer,
192            bytes_rx,
193            alive_rx,
194            writer_context,
195        ));
196
197        let subject = info.subject.clone();
198        let monitor_context = context;
199        // Spawn the connection monitor; errors are already logged inside
200        // wait_for_connection_tasks, so the Result is intentionally dropped.
201        tokio::spawn(async move {
202            let _ =
203                wait_for_connection_tasks(reader_task, writer_task, monitor_context, None, subject)
204                    .await;
205        });
206
207        // set up the prologue for the stream
208        // this might have transport specific metadata in the future
209        let prologue = Some(ResponseStreamPrologue {
210            error: None,
211            typed_error: None,
212        });
213
214        // create the stream sender
215        let stream_sender = StreamSender {
216            tx: bytes_tx,
217            prologue,
218        };
219
220        Ok(stream_sender)
221    }
222
223    /// Symmetric to [`Self::create_response_stream`] for the request-stream half:
224    /// dial the upstream TCP server with `StreamType::Request`, then return a
225    /// [`StreamReceiver`] that yields the data frames the upstream pushes down.
226    ///
227    /// The request stream is unidirectional after the handshake: the write half
228    /// is dropped as soon as the `CallHomeHandshake` is sent, so the downstream
229    /// never writes anything back (no `Sentinel` ack). The spawned reader task
230    /// forwards `TwoPartMessage::DataOnly` payloads into the channel and
231    /// translates `ControlMessage::Stop` / `Kill` into context cancellation;
232    /// `ControlMessage::Sentinel` terminates the task cleanly. A TCP close
233    /// before any `Sentinel` is treated as a truncated input (cancellation +
234    /// `context.kill()`), and dropping the returned `StreamReceiver` also stops
235    /// the task.
236    pub async fn create_request_stream(
237        context: Arc<dyn AsyncEngineContext>,
238        info: ConnectionInfo,
239        cancellation_counter: Option<IntCounter>,
240    ) -> Result<StreamReceiver> {
241        let info =
242            TcpStreamConnectionInfo::try_from(info).context("tcp-stream-connection-info-error")?;
243        tracing::trace!("Creating request stream for {:?}", info);
244
245        if info.stream_type != StreamType::Request {
246            return Err(error!(
247                "Invalid stream type; TcpClient::create_request_stream requires the stream type to be `request`; however {:?} was passed",
248                info.stream_type
249            ));
250        }
251
252        if info.context != context.id() {
253            return Err(error!(
254                "Invalid context; TcpClient::create_request_stream requires the context to be {:?}; however {:?} was passed",
255                context.id(),
256                info.context
257            ));
258        }
259
260        let (read_half, write_half) = TcpClient::connect_and_split(&info.address).await?;
261
262        let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
263        let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
264
265        let handshake = CallHomeHandshake {
266            subject: info.subject.clone(),
267            stream_type: StreamType::Request,
268        };
269        let handshake_bytes = serde_json::to_vec(&handshake).map_err(|err| {
270            error!(
271                "create_request_stream: Error converting CallHomeHandshake to JSON array: {err:#}"
272            )
273        })?;
274        framed_writer
275            .send(TwoPartMessage::from_header(handshake_bytes.into()))
276            .await
277            .map_err(|e| error!("failed to send request-stream handshake: {:?}", e))?;
278
279        // Request stream is unidirectional after the handshake: the downstream
280        // never writes again, so close the write half immediately.
281        drop(framed_writer);
282
283        let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel::<bytes::Bytes>(64);
284
285        let reader_span = tracing::Span::current();
286        tokio::spawn(
287            handle_request_reader(framed_reader, bytes_tx, context, cancellation_counter)
288                .instrument(reader_span),
289        );
290
291        Ok(StreamReceiver { rx: bytes_rx })
292    }
293}
294
295/// Cached TCP TLS connector — built once from env vars on first TCP connection, reused for every
296/// subsequent connection. `None` means plaintext; `Some(connector)` means TLS.
297///
298/// The connector itself is immutable after the first build (`rustls::ClientConfig` is cheaply
299/// `Arc`-cloned per-connection). The client *identity* it presents still hot-reloads from disk via
300/// `ReloadingCertifiedKey`, so rotating the client cert/key needs no restart; only the trust
301/// anchors (CA / insecure flag) are fixed at first build and require a restart to change.
302static TCP_TLS_CONNECTOR: once_cell::sync::OnceCell<Option<TlsConnector>> =
303    once_cell::sync::OnceCell::new();
304
305/// Return the process-wide `TlsConnector`, initialising it on first call.
306///
307/// Fails fast if any TLS env var is present but the configuration is incomplete
308/// (e.g. cert without key, or client cert without CA / insecure flag).
309fn get_tls_connector() -> anyhow::Result<&'static Option<TlsConnector>> {
310    TCP_TLS_CONNECTOR.get_or_try_init(build_tls_connector_from_env)
311}
312
313fn build_tls_connector_from_env() -> anyhow::Result<Option<TlsConnector>> {
314    use crate::config::environment_names::tcp_response_stream::tls as env;
315    let ca_cert_path = std::env::var(env::DYN_TCP_TLS_CA_CERT_PATH).ok();
316    let insecure = crate::config::env_is_truthy(env::DYN_TCP_TLS_INSECURE);
317    let client_cert = std::env::var(env::DYN_TCP_TLS_CLIENT_CERT_PATH).ok();
318    let client_key = std::env::var(env::DYN_TCP_TLS_CLIENT_KEY_PATH).ok();
319
320    let tls_requested =
321        ca_cert_path.is_some() || insecure || client_cert.is_some() || client_key.is_some();
322    if !tls_requested {
323        // Warn if the server side has TLS configured — a plaintext client connecting to a
324        // TLS server will fail with an opaque framing error rather than a clear TLS error.
325        let server_tls_set = std::env::var(env::DYN_TCP_TLS_CERT_PATH).is_ok();
326        if server_tls_set {
327            tracing::warn!(
328                "TCP client is running in plaintext mode but {} is set. \
329                 Set {} (or {} for dev) to enable client-side TLS.",
330                env::DYN_TCP_TLS_CERT_PATH,
331                env::DYN_TCP_TLS_CA_CERT_PATH,
332                env::DYN_TCP_TLS_INSECURE,
333            );
334        }
335        return Ok(None);
336    }
337    // Validate the client identity pair before the CA check so a lone cert/key
338    // produces the accurate "set both" error rather than a misleading CA error.
339    if client_cert.is_some() != client_key.is_some() {
340        anyhow::bail!(
341            "both {} and {} must be set together to present a client identity",
342            env::DYN_TCP_TLS_CLIENT_CERT_PATH,
343            env::DYN_TCP_TLS_CLIENT_KEY_PATH,
344        );
345    }
346    if !insecure && ca_cert_path.is_none() {
347        anyhow::bail!(
348            "TCP TLS is enabled but {} is not set and {} is not true; \
349             provide a CA cert or set insecure mode for development",
350            env::DYN_TCP_TLS_CA_CERT_PATH,
351            env::DYN_TCP_TLS_INSECURE,
352        );
353    }
354
355    let tls_config = crate::tls_utils::client_tls_config(
356        ca_cert_path.as_deref().map(std::path::Path::new),
357        insecure,
358        client_cert.as_deref().map(std::path::Path::new),
359        client_key.as_deref().map(std::path::Path::new),
360    )?;
361    Ok(Some(TlsConnector::from(Arc::new(tls_config))))
362}
363
364/// Extract the TLS `ServerName` for an outbound connection.
365///
366/// Checks `DYN_TCP_TLS_SERVER_NAME` first (explicit override), then falls back
367/// to the host part of `address`. Handles both IPv4 and IPv6 socket addresses.
368fn tls_server_name(address: &str) -> anyhow::Result<ServerName<'static>> {
369    use crate::config::environment_names::tcp_response_stream::tls as env;
370    let name = std::env::var(env::DYN_TCP_TLS_SERVER_NAME).unwrap_or_else(|_| {
371        // Parse as SocketAddr first — handles IPv6 like [::1]:443 correctly.
372        if let Ok(sock_addr) = address.parse::<std::net::SocketAddr>() {
373            sock_addr.ip().to_string()
374        } else {
375            // host:port string — strip the last colon-delimited segment.
376            match address.rfind(':') {
377                Some(pos) => address[..pos].to_owned(),
378                None => address.to_owned(),
379            }
380        }
381    });
382    ServerName::try_from(name).map_err(|e| anyhow::anyhow!("invalid TLS server name: {e}"))
383}
384
385async fn handle_request_reader(
386    mut framed_reader: FramedRead<BoxRead, TwoPartCodec>,
387    bytes_tx: tokio::sync::mpsc::Sender<bytes::Bytes>,
388    context: Arc<dyn AsyncEngineContext>,
389    cancellation_counter: Option<IntCounter>,
390) {
391    let cancellation_seen = {
392        // Keep cancellation and receiver-close notifications alive across frames instead of
393        // rebuilding their watch/Notify state for every request-stream message.
394        let killed = context.killed();
395        let stopped = context.stopped();
396        // The function explicitly drops `bytes_tx` after the loop, so let the persistent
397        // `closed()` future borrow a stream-lifetime clone instead.
398        let bytes_closed_tx = bytes_tx.clone();
399        let bytes_closed = bytes_closed_tx.closed();
400        tokio::pin!(killed, stopped, bytes_closed);
401
402        // Only mark cancellation on fatal errors or explicit upstream cancellation.
403        let mut cancellation_seen = false;
404        loop {
405            tokio::select! {
406                biased;
407
408                _ = &mut killed => {
409                    tracing::trace!("context kill signal received on request stream; shutting down");
410                    break;
411                }
412
413                _ = &mut stopped => {
414                    tracing::trace!("context stop signal received on request stream; shutting down");
415                    break;
416                }
417
418                // Downstream consumer dropped the StreamReceiver. Exit promptly
419                // instead of staying parked on `framed_reader.next()` until the
420                // socket closes — the data has nowhere to go. This is the consumer's
421                // own choice, so it is not a cancellation (no kill, no count).
422                _ = &mut bytes_closed => {
423                    tracing::debug!("downstream consumer dropped; exiting request-stream reader");
424                    break;
425                }
426
427                msg = framed_reader.next() => {
428                    match msg {
429                        Some(Ok(two_part_msg)) => match two_part_msg.into_message_type() {
430                            TwoPartMessageType::HeaderOnly(header) => {
431                                let ctrl = match serde_json::from_slice::<ControlMessage>(&header) {
432                                    Ok(c) => c,
433                                    Err(e) => {
434                                        tracing::warn!(
435                                            err = ?e,
436                                            "invalid control message, closing connection"
437                                        );
438                                        cancellation_seen = true;
439                                        context.kill();
440                                        break;
441                                    }
442                                };
443                                match ctrl {
444                                    ControlMessage::Stop => {
445                                        cancellation_seen = true;
446                                        record_upstream_cancellation(context.as_ref(), "stop");
447                                        context.stop();
448                                        break;
449                                    }
450                                    ControlMessage::Kill => {
451                                        cancellation_seen = true;
452                                        record_upstream_cancellation(context.as_ref(), "kill");
453                                        context.kill();
454                                        break;
455                                    }
456                                    ControlMessage::Sentinel => {
457                                        tracing::trace!("upstream signaled end of request stream");
458                                        break;
459                                    }
460                                }
461                            }
462                            TwoPartMessageType::DataOnly(data) => {
463                                if bytes_tx.send(data).await.is_err() {
464                                    tracing::debug!("downstream consumer dropped; exiting request-stream reader");
465                                    break;
466                                }
467                            }
468                            _ => {
469                                tracing::warn!("fatal error - unexpected message shape on request stream");
470                                cancellation_seen = true;
471                                context.kill();
472                                break;
473                            }
474                        }
475                        Some(Err(e)) => {
476                            tracing::warn!("fatal error - failed to decode message on request stream: {e:?}");
477                            cancellation_seen = true;
478                            context.kill();
479                            break;
480                        }
481                        None => {
482                            // Socket closed before a Sentinel/Stop/Kill: the request
483                            // input is truncated. Kill the context so the consumer
484                            // sees an aborted stream rather than a clean end, and
485                            // count it as a cancellation.
486                            tracing::warn!("request stream closed by upstream before sentinel; treating as truncated");
487                            cancellation_seen = true;
488                            context.kill();
489                            break;
490                        }
491                    }
492                }
493            }
494        }
495
496        // Leaving this scope drops the pinned `closed()` future and its sender clone
497        // before the original sender below signals end-of-stream.
498        cancellation_seen
499    };
500
501    if cancellation_seen && let Some(counter) = &cancellation_counter {
502        counter.inc();
503    }
504
505    // Dropping bytes_tx closes the receiver side, signaling end-of-stream to the
506    // engine consumer.
507    drop(bytes_tx);
508}
509
510async fn wait_for_connection_tasks(
511    reader_task: tokio::task::JoinHandle<FramedRead<BoxRead, TwoPartCodec>>,
512    writer_task: tokio::task::JoinHandle<Result<FramedWrite<BoxWrite, TwoPartCodec>>>,
513    context: Arc<dyn AsyncEngineContext>,
514    peer_port: Option<u16>,
515    subject: String,
516) -> Result<()> {
517    // Await the reader first and abort the writer on reader Err — the
518    // writer parks on `bytes_rx.recv()` and won't wake on its own.
519    let reader = match reader_task.await {
520        Ok(reader) => reader,
521        Err(reader_err) => {
522            writer_task.abort();
523            let _ = writer_task.await;
524            tracing::error!(
525                subject = %subject,
526                peer_port = ?peer_port,
527                err = ?reader_err,
528                "reader task failed to join"
529            );
530            return Err(reader_err.into());
531        }
532    };
533
534    match writer_task.await {
535        Ok(Ok(_)) => {}
536        Ok(Err(e)) => {
537            tracing::error!(
538                subject = %subject,
539                peer_port = ?peer_port,
540                err = ?e,
541                "writer task returned error"
542            );
543            return Err(e);
544        }
545        Err(writer_err) => {
546            tracing::error!(
547                subject = %subject,
548                peer_port = ?peer_port,
549                err = ?writer_err,
550                "writer task failed to join"
551            );
552            return Err(writer_err.into());
553        }
554    }
555
556    // Drain the read half until the server closes the connection (FIN).
557    // The write half is dropped above when the FramedWrite is discarded.
558    let read_half = reader.into_inner();
559    wait_for_server_shutdown(read_half, context).await
560}
561
562async fn wait_for_server_shutdown(
563    mut reader: BoxRead,
564    context: Arc<dyn AsyncEngineContext>,
565) -> Result<()> {
566    // `handle_writer` skips the closing sentinel on both `killed` and
567    // `stopped`, so the server has nothing to react to in either case;
568    // sitting in the read loop until the 10 s deadline would be dead time.
569    if context.is_killed() || context.is_stopped() {
570        tracing::debug!("stream context killed or stopped; skipping server FIN wait");
571        return Ok(());
572    }
573
574    // Await the tcp server to shutdown the socket connection, bounded by a
575    // timeout so normal sentinel shutdown cannot hang indefinitely.
576    let mut buf = [0u8; 1024];
577    let deadline = Instant::now() + Duration::from_secs(10);
578    loop {
579        let n = time::timeout_at(deadline, reader.read(&mut buf))
580            .await
581            .inspect_err(|_| {
582                tracing::debug!("server did not close socket within the deadline");
583            })?
584            .inspect_err(|e| {
585                tracing::debug!(err = ?e, "failed to read from stream");
586            })?;
587        if n == 0 {
588            // Server has closed (FIN)
589            break;
590        }
591    }
592
593    Ok(())
594}
595
596async fn handle_reader(
597    framed_reader: FramedRead<BoxRead, TwoPartCodec>,
598    context: Arc<dyn AsyncEngineContext>,
599    alive_tx: tokio::sync::oneshot::Sender<()>,
600    cancellation_counter: Option<IntCounter>,
601) -> FramedRead<BoxRead, TwoPartCodec> {
602    let mut framed_reader = framed_reader;
603    let mut alive_tx = alive_tx;
604    // Set on every cancellation arm; counted once after the loop.
605    let mut cancellation_seen = false;
606    loop {
607        tokio::select! {
608            msg = framed_reader.next() => {
609                match msg {
610                    Some(Ok(two_part_msg)) => {
611                        match two_part_msg.optional_parts() {
612                           (Some(bytes), None) => {
613                                let msg = match serde_json::from_slice::<ControlMessage>(bytes) {
614                                    Ok(msg) => msg,
615                                    Err(e) => {
616                                        tracing::warn!(
617                                            err = ?e,
618                                            "invalid control message, closing connection"
619                                        );
620                                        cancellation_seen = true;
621                                        context.kill();
622                                        break;
623                                    }
624                                };
625
626                                // Stop/Kill intentionally do not `break`: the
627                                // reader keeps running so a later Kill can
628                                // upgrade an earlier Stop (and vice versa).
629                                // The loop still exits promptly via the
630                                // `alive_tx.closed()` arm once `handle_writer`
631                                // reacts to `context.stop()` / `context.kill()`.
632                                match msg {
633                                    ControlMessage::Stop => {
634                                        cancellation_seen = true;
635                                        record_upstream_cancellation(context.as_ref(), "stop");
636                                        context.stop();
637                                    }
638                                    ControlMessage::Kill => {
639                                        cancellation_seen = true;
640                                        record_upstream_cancellation(context.as_ref(), "kill");
641                                        context.kill();
642                                    }
643                                    ControlMessage::Sentinel => {
644                                        tracing::warn!(
645                                            "unexpected sentinel on client reader, closing connection"
646                                        );
647                                        cancellation_seen = true;
648                                        context.kill();
649                                        break;
650                                    }
651                                }
652                           }
653                           _ => {
654                                tracing::warn!(
655                                    "unexpected non-control message on client reader, closing connection"
656                                );
657                                cancellation_seen = true;
658                                context.kill();
659                                break;
660                           }
661                        }
662                    }
663                    Some(Err(e)) => {
664                        // Kill the engine context so the producer stops
665                        // generating responses that can no longer be delivered.
666                        tracing::warn!(err = ?e, "tcp stream read error, closing connection");
667                        cancellation_seen = true;
668                        context.kill();
669                        break;
670                    }
671                    None => {
672                        tracing::debug!("tcp stream closed by server");
673                        break;
674                    }
675                }
676            }
677            _ = alive_tx.closed() => {
678                break;
679            }
680        }
681    }
682    if cancellation_seen && let Some(counter) = &cancellation_counter {
683        counter.inc();
684    }
685    framed_reader
686}
687
688async fn handle_writer(
689    mut framed_writer: FramedWrite<BoxWrite, TwoPartCodec>,
690    mut bytes_rx: tokio::sync::mpsc::Receiver<TwoPartMessage>,
691    alive_rx: tokio::sync::oneshot::Receiver<()>,
692    context: Arc<dyn AsyncEngineContext>,
693) -> Result<FramedWrite<BoxWrite, TwoPartCodec>> {
694    // Keep one cancellation future per stream. Recreating these futures for every queued
695    // frame repeatedly clones the context's watch receivers and churns Notify state.
696    let killed = context.killed();
697    let stopped = context.stopped();
698    tokio::pin!(killed, stopped);
699
700    // Only send sentinel for normal channel closure
701    let mut send_sentinel = true;
702
703    loop {
704        let msg = tokio::select! {
705            biased;
706
707            _ = &mut killed => {
708                tracing::trace!("context kill signal received; shutting down");
709                send_sentinel = false;
710                break;
711            }
712
713            _ = &mut stopped => {
714                tracing::trace!("context stop signal received; shutting down");
715                send_sentinel = false;
716                break;
717            }
718
719            msg = bytes_rx.recv() => {
720                match msg {
721                    Some(msg) => msg,
722                    None => {
723                        tracing::trace!("response channel closed; shutting down");
724                        break;
725                    }
726                }
727            }
728        };
729
730        if let Err(e) = framed_writer.send(msg).await {
731            tracing::trace!(
732                "failed to send message to network; possible disconnect: {:?}",
733                e
734            );
735            send_sentinel = false;
736            break;
737        }
738    }
739
740    // Send sentinel only on normal closure
741    if send_sentinel {
742        let message = serde_json::to_vec(&ControlMessage::Sentinel)?;
743        let msg = TwoPartMessage::from_header(message.into());
744        framed_writer.send(msg).await?;
745    }
746
747    drop(alive_rx);
748    Ok(framed_writer)
749}
750
751#[cfg(test)]
752mod tests {
753    use super::*;
754    use crate::pipeline::context::Controller;
755    use crate::pipeline::network::tcp::test_utils::create_tcp_pair;
756    use bytes::Bytes;
757    use futures::StreamExt;
758    use std::collections::HashMap;
759    use std::sync::{Arc, Mutex};
760    use tokio::io::{AsyncReadExt, AsyncWriteExt};
761    use tokio::net::TcpStream;
762    use tokio::sync::{mpsc, oneshot};
763    use tokio_util::codec::FramedRead;
764    use tracing::field::{Field, Visit};
765    use tracing_subscriber::Layer;
766    use tracing_subscriber::layer::{Context as TraceContext, SubscriberExt};
767    use tracing_subscriber::registry::LookupSpan;
768    use tracing_subscriber::util::SubscriberInitExt;
769
770    type CapturedCancellationEvent = (HashMap<String, String>, Option<String>);
771
772    #[derive(Default)]
773    struct CancellationEventCapture(Mutex<Vec<CapturedCancellationEvent>>);
774
775    struct EventFieldVisitor<'a>(&'a mut HashMap<String, String>);
776
777    impl Visit for EventFieldVisitor<'_> {
778        fn record_str(&mut self, field: &Field, value: &str) {
779            self.0.insert(field.name().to_string(), value.to_string());
780        }
781
782        fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
783            self.0
784                .insert(field.name().to_string(), format!("{value:?}"));
785        }
786    }
787
788    struct CancellationEventLayer(Arc<CancellationEventCapture>);
789
790    impl<S> Layer<S> for CancellationEventLayer
791    where
792        S: tracing::Subscriber + for<'lookup> LookupSpan<'lookup>,
793    {
794        fn on_event(&self, event: &tracing::Event<'_>, ctx: TraceContext<'_, S>) {
795            if event.metadata().target() != "request_span" {
796                return;
797            }
798            let mut fields = HashMap::new();
799            event.record(&mut EventFieldVisitor(&mut fields));
800            if fields
801                .get("message")
802                .is_none_or(|message| message.trim_matches('"') != "request cancellation received")
803            {
804                return;
805            }
806            let parent = ctx.event_span(event).map(|span| span.name().to_string());
807            self.0.0.lock().unwrap().push((fields, parent));
808        }
809    }
810
811    #[tokio::test]
812    async fn upstream_cancellation_event_is_parented_to_worker_request_span() {
813        let captured = Arc::new(CancellationEventCapture::default());
814        let _subscriber = tracing_subscriber::registry()
815            .with(CancellationEventLayer(captured.clone()))
816            .set_default();
817        let controller = Arc::new(Controller::new("request-123".to_string()));
818        let span = tracing::info_span!(target: "request_span", "handle_payload");
819
820        async {
821            record_upstream_cancellation(controller.as_ref(), "stop");
822        }
823        .instrument(span)
824        .await;
825
826        let events = captured.0.lock().unwrap();
827        assert_eq!(events.len(), 1);
828        assert_eq!(
829            events[0].0.get("cancellation.signal").map(String::as_str),
830            Some("stop")
831        );
832        assert_eq!(
833            events[0].0.get("cancellation.source").map(String::as_str),
834            Some("upstream")
835        );
836        assert_eq!(events[0].1.as_deref(), Some("handle_payload"));
837    }
838
839    struct WriterHarness {
840        server: tokio::net::TcpStream,
841        framed_writer: FramedWrite<BoxWrite, TwoPartCodec>,
842        bytes_tx: mpsc::Sender<TwoPartMessage>,
843        bytes_rx: mpsc::Receiver<TwoPartMessage>,
844        alive_tx: oneshot::Sender<()>,
845        alive_rx: oneshot::Receiver<()>,
846        controller: Arc<Controller>,
847    }
848
849    /// Creates a reusable writer harness with paired TCP streams and test channels.
850    async fn writer_harness() -> WriterHarness {
851        let (client, server) = create_tcp_pair().await;
852        let (_, write_half) = tokio::io::split(client);
853        let framed_writer =
854            FramedWrite::new(Box::new(write_half) as BoxWrite, TwoPartCodec::default());
855
856        let (bytes_tx, bytes_rx) = mpsc::channel(64);
857        let (alive_tx, alive_rx) = oneshot::channel::<()>();
858        let controller = Arc::new(Controller::default());
859
860        WriterHarness {
861            server,
862            framed_writer,
863            bytes_tx,
864            bytes_rx,
865            alive_tx,
866            alive_rx,
867            controller,
868        }
869    }
870
871    async fn recv_msg(reader: &mut FramedRead<TcpStream, TwoPartCodec>) -> TwoPartMessage {
872        reader
873            .next()
874            .await
875            .expect("expected message")
876            .expect("failed to decode message")
877    }
878
879    fn assert_data_only_message(msg: TwoPartMessage, expected: &[u8]) {
880        let (header, data) = msg.optional_parts();
881        assert!(header.is_none(), "data-only message should not have header");
882        assert_eq!(
883            data.expect("data payload missing").as_ref(),
884            expected,
885            "data payload should match"
886        );
887    }
888
889    fn assert_header_only_message(msg: TwoPartMessage, expected: &[u8]) {
890        let (header, data) = msg.optional_parts();
891        assert!(data.is_none(), "header-only message should not carry data");
892        assert_eq!(
893            header.expect("header missing").as_ref(),
894            expected,
895            "header payload should match"
896        );
897    }
898
899    fn assert_header_and_data_message(
900        msg: TwoPartMessage,
901        expected_header: &[u8],
902        expected_data: &[u8],
903    ) {
904        let (header, data) = msg.optional_parts();
905        assert_eq!(
906            header.expect("header missing").as_ref(),
907            expected_header,
908            "header payload should match"
909        );
910        assert_eq!(
911            data.expect("data missing").as_ref(),
912            expected_data,
913            "data payload should match"
914        );
915    }
916
917    fn assert_sentinel_message(msg: TwoPartMessage) {
918        let (header, data) = msg.optional_parts();
919        assert!(data.is_none(), "sentinel should not include a data section");
920        let expected_sentinel = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
921        assert_eq!(
922            header.expect("sentinel header missing").as_ref(),
923            expected_sentinel.as_slice(),
924            "sentinel header should match serialized ControlMessage::Sentinel"
925        );
926    }
927
928    /// Test that handle_writer forwards messages from the channel to the framed writer
929    #[tokio::test]
930    async fn test_handle_writer_forwards_messages() {
931        let WriterHarness {
932            server,
933            framed_writer,
934            bytes_tx,
935            bytes_rx,
936            alive_rx,
937            controller,
938            ..
939        } = writer_harness().await;
940
941        // Send test messages
942        let test_msg = TwoPartMessage::from_data(Bytes::from("test data"));
943        bytes_tx.send(test_msg).await.unwrap();
944
945        // Close the sender to trigger normal termination
946        drop(bytes_tx);
947
948        let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
949
950        assert!(result.is_ok());
951
952        // Decode from server side to verify data and sentinel were sent
953        let mut reader = FramedRead::new(server, TwoPartCodec::default());
954
955        let msg = recv_msg(&mut reader).await;
956        assert_data_only_message(msg, b"test data");
957
958        let sentinel = recv_msg(&mut reader).await;
959        assert_sentinel_message(sentinel);
960    }
961
962    /// Test that handle_writer sends sentinel on normal channel closure
963    #[tokio::test]
964    async fn test_handle_writer_sends_sentinel_on_normal_closure() {
965        let WriterHarness {
966            mut server,
967            framed_writer,
968            bytes_tx,
969            bytes_rx,
970            alive_rx,
971            controller,
972            ..
973        } = writer_harness().await;
974
975        // Close the sender immediately to trigger normal termination
976        drop(bytes_tx);
977
978        let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
979
980        assert!(result.is_ok());
981
982        // Read from server side to verify sentinel was sent
983        let mut buffer = vec![0u8; 1024];
984        let n = server.read(&mut buffer).await.unwrap();
985
986        // Buffer should contain the sentinel message
987        assert!(n > 0, "Expected sentinel to be written to the TCP stream");
988
989        // Verify it contains the sentinel message by checking for the JSON
990        let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
991        assert!(
992            buffer[..n]
993                .windows(sentinel_json.len())
994                .any(|w| w == sentinel_json.as_slice()),
995            "Buffer should contain sentinel message. Buffer: {:?}",
996            String::from_utf8_lossy(&buffer[..n])
997        );
998    }
999
1000    /// A normally completed response stream sends Sentinel, receives the
1001    /// server's FIN, and leaves the cancellation metric unchanged.
1002    #[tokio::test]
1003    async fn test_normal_response_stream_completion_does_not_count_cancellation() {
1004        let (client, server) = create_tcp_pair().await;
1005        let (read_half, write_half) = tokio::io::split(client);
1006        let framed_reader =
1007            FramedRead::new(Box::new(read_half) as BoxRead, TwoPartCodec::default());
1008        let framed_writer =
1009            FramedWrite::new(Box::new(write_half) as BoxWrite, TwoPartCodec::default());
1010        let (bytes_tx, bytes_rx) = mpsc::channel(64);
1011        let (alive_tx, alive_rx) = oneshot::channel::<()>();
1012        let controller = Arc::new(Controller::default());
1013        let cancellation_counter = IntCounter::new(
1014            "tcp_client_normal_completion_cancellations_test",
1015            "test cancellation counter",
1016        )
1017        .unwrap();
1018
1019        let reader_context = controller.clone();
1020        let counter_clone = cancellation_counter.clone();
1021        let reader_task = tokio::spawn(async move {
1022            handle_reader(framed_reader, reader_context, alive_tx, Some(counter_clone)).await
1023        });
1024        let writer_context = controller.clone();
1025        let writer_task = tokio::spawn(async move {
1026            handle_writer(framed_writer, bytes_rx, alive_rx, writer_context).await
1027        });
1028
1029        drop(bytes_tx);
1030
1031        let mut server_reader = FramedRead::new(server, TwoPartCodec::default());
1032        let sentinel = recv_msg(&mut server_reader).await;
1033        assert_sentinel_message(sentinel);
1034        drop(server_reader);
1035
1036        writer_task.await.unwrap().unwrap();
1037        reader_task.await.unwrap();
1038
1039        assert!(
1040            !controller.is_stopped() && !controller.is_killed(),
1041            "normal response completion must not cancel the context"
1042        );
1043        assert_eq!(
1044            cancellation_counter.get(),
1045            0,
1046            "normal response completion must not increment the cancellation counter"
1047        );
1048    }
1049
1050    /// Test that handle_writer does NOT send sentinel when context is killed
1051    #[tokio::test]
1052    async fn test_handle_writer_no_sentinel_on_context_killed() {
1053        let WriterHarness {
1054            mut server,
1055            framed_writer,
1056            bytes_rx,
1057            alive_rx,
1058            controller,
1059            ..
1060        } = writer_harness().await;
1061
1062        // Kill the context
1063        controller.kill();
1064
1065        let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
1066
1067        assert!(result.is_ok());
1068
1069        // Drop the writer to close the connection, then try to read. Otherwise,
1070        // the test will hang on `server.read()`
1071        drop(result);
1072
1073        // Read from server side - should get no sentinel
1074        let mut buffer = vec![0u8; 1024];
1075        let n = server.read(&mut buffer).await.unwrap();
1076
1077        // Buffer should be empty (no sentinel sent)
1078        let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
1079        assert!(
1080            n == 0
1081                || !buffer[..n]
1082                    .windows(sentinel_json.len())
1083                    .any(|w| w == sentinel_json.as_slice()),
1084            "Buffer should NOT contain sentinel message when context is killed"
1085        );
1086    }
1087
1088    /// Test that handle_writer does NOT send sentinel when context is stopped
1089    #[tokio::test]
1090    async fn test_handle_writer_no_sentinel_on_context_stopped() {
1091        let WriterHarness {
1092            mut server,
1093            framed_writer,
1094            bytes_rx,
1095            alive_rx,
1096            controller,
1097            ..
1098        } = writer_harness().await;
1099
1100        // Stop the context
1101        controller.stop();
1102
1103        let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
1104
1105        assert!(result.is_ok());
1106
1107        // Drop the writer to close the connection, then try to read. Otherwise,
1108        // the test will hang on `server.read()`
1109        drop(result);
1110
1111        // Read from server side - should get no sentinel
1112        let mut buffer = vec![0u8; 1024];
1113        let n = server.read(&mut buffer).await.unwrap();
1114
1115        // Buffer should be empty (no sentinel sent)
1116        let sentinel_json = serde_json::to_vec(&ControlMessage::Sentinel).unwrap();
1117        assert!(
1118            n == 0
1119                || !buffer[..n]
1120                    .windows(sentinel_json.len())
1121                    .any(|w| w == sentinel_json.as_slice()),
1122            "Buffer should NOT contain sentinel message when context is stopped"
1123        );
1124    }
1125
1126    /// Test that handle_writer handles multiple messages correctly
1127    #[tokio::test]
1128    async fn test_handle_writer_multiple_messages() {
1129        let WriterHarness {
1130            server,
1131            framed_writer,
1132            bytes_tx,
1133            bytes_rx,
1134            alive_rx,
1135            controller,
1136            ..
1137        } = writer_harness().await;
1138
1139        // Send multiple messages
1140        for i in 0..5 {
1141            let test_msg = TwoPartMessage::from_data(Bytes::from(format!("message {}", i)));
1142            bytes_tx.send(test_msg).await.unwrap();
1143        }
1144
1145        // Close the sender to trigger normal termination
1146        drop(bytes_tx);
1147
1148        let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
1149
1150        assert!(result.is_ok());
1151
1152        // Decode from server side to verify all messages plus sentinel
1153        let mut reader = FramedRead::new(server, TwoPartCodec::default());
1154        for i in 0..5 {
1155            let msg = recv_msg(&mut reader).await;
1156            assert_data_only_message(msg, format!("message {}", i).as_bytes());
1157        }
1158
1159        let sentinel = recv_msg(&mut reader).await;
1160        assert_sentinel_message(sentinel);
1161    }
1162
1163    /// Test that alive_rx is dropped after handle_writer completes
1164    #[tokio::test]
1165    async fn test_handle_writer_drops_alive_rx() {
1166        let WriterHarness {
1167            framed_writer,
1168            bytes_tx,
1169            bytes_rx,
1170            alive_tx,
1171            alive_rx,
1172            controller,
1173            ..
1174        } = writer_harness().await;
1175
1176        // Close the sender to trigger normal termination
1177        drop(bytes_tx);
1178
1179        let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
1180
1181        assert!(result.is_ok());
1182
1183        // alive_tx should now be closed because alive_rx was dropped
1184        assert!(alive_tx.is_closed());
1185    }
1186
1187    /// Test handle_writer with header-only messages (control messages)
1188    #[tokio::test]
1189    async fn test_handle_writer_header_only_messages() {
1190        let WriterHarness {
1191            server,
1192            framed_writer,
1193            bytes_tx,
1194            bytes_rx,
1195            alive_rx,
1196            controller,
1197            ..
1198        } = writer_harness().await;
1199
1200        // Send a header-only message
1201        let header_msg = TwoPartMessage::from_header(Bytes::from("header content"));
1202        bytes_tx.send(header_msg).await.unwrap();
1203
1204        // Close the sender
1205        drop(bytes_tx);
1206
1207        let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
1208
1209        assert!(result.is_ok());
1210
1211        let mut reader = FramedRead::new(server, TwoPartCodec::default());
1212
1213        let header_msg = recv_msg(&mut reader).await;
1214        assert_header_only_message(header_msg, b"header content");
1215
1216        let sentinel = recv_msg(&mut reader).await;
1217        assert_sentinel_message(sentinel);
1218    }
1219
1220    /// Test handle_writer with mixed header and data messages
1221    #[tokio::test]
1222    async fn test_handle_writer_mixed_messages() {
1223        let WriterHarness {
1224            server,
1225            framed_writer,
1226            bytes_tx,
1227            bytes_rx,
1228            alive_rx,
1229            controller,
1230            ..
1231        } = writer_harness().await;
1232
1233        // Send mixed messages
1234        bytes_tx
1235            .send(TwoPartMessage::from_header(Bytes::from("header1")))
1236            .await
1237            .unwrap();
1238        bytes_tx
1239            .send(TwoPartMessage::from_data(Bytes::from("data1")))
1240            .await
1241            .unwrap();
1242        bytes_tx
1243            .send(TwoPartMessage::from_parts(
1244                Bytes::from("header2"),
1245                Bytes::from("data2"),
1246            ))
1247            .await
1248            .unwrap();
1249
1250        // Close the sender
1251        drop(bytes_tx);
1252
1253        let result = handle_writer(framed_writer, bytes_rx, alive_rx, controller).await;
1254
1255        assert!(result.is_ok());
1256
1257        let mut reader = FramedRead::new(server, TwoPartCodec::default());
1258
1259        let first = recv_msg(&mut reader).await;
1260        assert_header_only_message(first, b"header1");
1261
1262        let second = recv_msg(&mut reader).await;
1263        assert_data_only_message(second, b"data1");
1264
1265        let third = recv_msg(&mut reader).await;
1266        assert_header_and_data_message(third, b"header2", b"data2");
1267
1268        let sentinel = recv_msg(&mut reader).await;
1269        assert_sentinel_message(sentinel);
1270    }
1271
1272    /// Killed or stopped contexts skip the server FIN deadline.
1273    #[tokio::test]
1274    async fn test_wait_for_server_shutdown_skips_terminal_context() {
1275        for action in [Controller::kill as fn(&Controller), Controller::stop] {
1276            let (client, _server) = create_tcp_pair().await;
1277            let controller = Arc::new(Controller::default());
1278            action(&controller);
1279
1280            let context: Arc<dyn AsyncEngineContext> = controller;
1281            let result = tokio::time::timeout(
1282                std::time::Duration::from_millis(50),
1283                wait_for_server_shutdown(Box::new(client), context),
1284            )
1285            .await;
1286
1287            assert!(result.is_ok(), "terminal context should not wait for FIN");
1288            assert!(
1289                result.unwrap().is_ok(),
1290                "terminal context shutdown should succeed"
1291            );
1292        }
1293    }
1294
1295    /// Read error in the connection monitor kills the context and skips the FIN wait.
1296    #[tokio::test]
1297    async fn test_connection_monitor_skips_fin_wait_after_read_error_kills_context() {
1298        let (client, mut server) = create_tcp_pair().await;
1299        let (read_half, write_half) = tokio::io::split(client);
1300        let framed_reader =
1301            FramedRead::new(Box::new(read_half) as BoxRead, TwoPartCodec::default());
1302        let framed_writer =
1303            FramedWrite::new(Box::new(write_half) as BoxWrite, TwoPartCodec::default());
1304        let (_bytes_tx, bytes_rx) = mpsc::channel(64);
1305        let (alive_tx, alive_rx) = oneshot::channel::<()>();
1306        let controller = Arc::new(Controller::default());
1307
1308        let reader_context = controller.clone();
1309        let reader_task = tokio::spawn(async move {
1310            handle_reader(framed_reader, reader_context, alive_tx, None).await
1311        });
1312        let writer_context = controller.clone();
1313        let writer_task = tokio::spawn(async move {
1314            handle_writer(framed_writer, bytes_rx, alive_rx, writer_context).await
1315        });
1316
1317        // Bypass the codec and write a complete but invalid TwoPartCodec
1318        // header. This drives the client reader into Some(Err(_)) without
1319        // closing the server side of the socket.
1320        server.write_all(&[0xFF; 24]).await.unwrap();
1321
1322        let monitor_context: Arc<dyn AsyncEngineContext> = controller.clone();
1323        let result = tokio::time::timeout(
1324            std::time::Duration::from_millis(250),
1325            wait_for_connection_tasks(
1326                reader_task,
1327                writer_task,
1328                monitor_context,
1329                None,
1330                "test-subject".to_string(),
1331            ),
1332        )
1333        .await;
1334
1335        assert!(
1336            result.is_ok(),
1337            "connection monitor should not wait for the FIN deadline after read error"
1338        );
1339        assert!(result.unwrap().is_ok(), "connection monitor should succeed");
1340        assert!(
1341            controller.is_killed(),
1342            "read error should kill the stream context"
1343        );
1344    }
1345
1346    /// Reader-side panic must abort the writer and return promptly rather than
1347    /// hanging on `tokio::join!`. Locks in the fix added with this function's
1348    /// sequential-await + writer-abort behavior.
1349    ///
1350    /// Setup: spawn a reader task that panics immediately (so
1351    /// `reader_task.await` yields `Err(JoinError::panic)`), and a writer task
1352    /// that parks indefinitely waiting for application bytes (so without the
1353    /// abort, `tokio::join!` on the previous implementation would never wake).
1354    /// Expect: `wait_for_connection_tasks` returns Err within the timeout.
1355    #[tokio::test]
1356    async fn test_connection_monitor_aborts_writer_when_reader_panics() {
1357        // Reader task that panics immediately. The explicit JoinHandle type
1358        // pins the inferred return type to the one wait_for_connection_tasks
1359        // expects; `panic!` is type `!`, which coerces to that type.
1360        let reader_task: tokio::task::JoinHandle<FramedRead<BoxRead, TwoPartCodec>> =
1361            tokio::spawn(async {
1362                panic!("simulated reader panic to trigger JoinError");
1363            });
1364
1365        // Writer task that would block indefinitely waiting on application
1366        // bytes. Under the pre-fix `tokio::join!` implementation, this would
1367        // prevent the function from returning when the reader panicked.
1368        // After the fix, the abort drives this task to completion promptly.
1369        let writer_task: tokio::task::JoinHandle<Result<FramedWrite<BoxWrite, TwoPartCodec>>> =
1370            tokio::spawn(async {
1371                std::future::pending::<()>().await;
1372                unreachable!()
1373            });
1374
1375        let controller = Arc::new(Controller::default());
1376        let context: Arc<dyn AsyncEngineContext> = controller.clone();
1377
1378        // 250 ms is generous — the abort + JoinHandle resolution should fire
1379        // sub-millisecond. We are checking for "doesn't hang", not "fast".
1380        let result = tokio::time::timeout(
1381            std::time::Duration::from_millis(250),
1382            wait_for_connection_tasks(
1383                reader_task,
1384                writer_task,
1385                context,
1386                None,
1387                "test-reader-panic".to_string(),
1388            ),
1389        )
1390        .await;
1391
1392        // Outer timeout must not fire: the abort path must surface the reader
1393        // JoinError before the writer would have produced any bytes.
1394        assert!(
1395            result.is_ok(),
1396            "wait_for_connection_tasks must return after reader panic, \
1397             not hang waiting on the writer"
1398        );
1399
1400        // The inner result must be Err — the reader's JoinError propagates.
1401        assert!(
1402            result.unwrap().is_err(),
1403            "reader panic should propagate as Err from wait_for_connection_tasks"
1404        );
1405    }
1406
1407    // ==================== handle_reader tests ====================
1408
1409    struct ReaderHarness {
1410        framed_server: FramedWrite<BoxWrite, TwoPartCodec>,
1411        framed_reader: FramedRead<BoxRead, TwoPartCodec>,
1412        alive_tx: oneshot::Sender<()>,
1413        alive_rx: oneshot::Receiver<()>,
1414        controller: Arc<Controller>,
1415    }
1416
1417    /// Creates a reusable reader harness with paired TCP streams and test channels.
1418    async fn reader_harness() -> ReaderHarness {
1419        let (client, server) = create_tcp_pair().await;
1420        let (read_half, _write_half) = tokio::io::split(client);
1421        let (_server_read, server_write) = tokio::io::split(server);
1422
1423        let framed_reader =
1424            FramedRead::new(Box::new(read_half) as BoxRead, TwoPartCodec::default());
1425        let framed_server =
1426            FramedWrite::new(Box::new(server_write) as BoxWrite, TwoPartCodec::default());
1427        let (alive_tx, alive_rx) = oneshot::channel::<()>();
1428        let controller = Arc::new(Controller::default());
1429
1430        ReaderHarness {
1431            framed_server,
1432            framed_reader,
1433            alive_tx,
1434            alive_rx,
1435            controller,
1436        }
1437    }
1438
1439    fn control_message(msg: &ControlMessage) -> TwoPartMessage {
1440        let msg_bytes = serde_json::to_vec(msg).unwrap();
1441        TwoPartMessage::from_header(Bytes::from(msg_bytes))
1442    }
1443
1444    /// Test that handle_reader handles Stop control message by calling context.stop()
1445    #[tokio::test]
1446    async fn test_handle_reader_stop_control_message() {
1447        let ReaderHarness {
1448            mut framed_server,
1449            framed_reader,
1450            alive_tx,
1451            alive_rx: _alive_rx,
1452            controller,
1453        } = reader_harness().await;
1454
1455        // Spawn the reader task
1456        let controller_clone = controller.clone();
1457        let reader_handle = tokio::spawn(async move {
1458            handle_reader(framed_reader, controller_clone, alive_tx, None).await
1459        });
1460
1461        // Send Stop control message from server
1462        framed_server
1463            .send(control_message(&ControlMessage::Stop))
1464            .await
1465            .unwrap();
1466
1467        // Close the framed server to signal EOF to the client
1468        framed_server.close().await.unwrap();
1469
1470        // Wait for reader to finish
1471        let _ = reader_handle.await.unwrap();
1472
1473        // Verify that stop was called on the controller
1474        assert!(
1475            controller.is_stopped(),
1476            "Controller should be stopped after receiving Stop message"
1477        );
1478    }
1479
1480    /// Test that handle_reader handles Kill control message by calling context.kill()
1481    #[tokio::test]
1482    async fn test_handle_reader_kill_control_message() {
1483        let ReaderHarness {
1484            mut framed_server,
1485            framed_reader,
1486            alive_tx,
1487            alive_rx: _alive_rx,
1488            controller,
1489        } = reader_harness().await;
1490
1491        // Spawn the reader task
1492        let controller_clone = controller.clone();
1493        let reader_handle = tokio::spawn(async move {
1494            handle_reader(framed_reader, controller_clone, alive_tx, None).await
1495        });
1496
1497        // Send Kill control message from server
1498        framed_server
1499            .send(control_message(&ControlMessage::Kill))
1500            .await
1501            .unwrap();
1502
1503        // Close the framed server to signal EOF to the client
1504        framed_server.close().await.unwrap();
1505
1506        // Wait for reader to finish
1507        let _ = reader_handle.await.unwrap();
1508
1509        // Verify that kill was called on the controller
1510        assert!(
1511            controller.is_killed(),
1512            "Controller should be killed after receiving Kill message"
1513        );
1514    }
1515
1516    /// Test that handle_reader exits when alive channel is closed
1517    #[tokio::test]
1518    async fn test_handle_reader_exits_on_alive_channel_closed() {
1519        let ReaderHarness {
1520            framed_reader,
1521            alive_tx,
1522            alive_rx,
1523            controller,
1524            ..
1525        } = reader_harness().await;
1526
1527        // Spawn the reader task
1528        let reader_handle =
1529            tokio::spawn(
1530                async move { handle_reader(framed_reader, controller, alive_tx, None).await },
1531            );
1532
1533        // Drop the alive_rx to close the channel (simulating writer finishing)
1534        drop(alive_rx);
1535
1536        // Reader should exit due to alive channel closure
1537        let result = reader_handle.await;
1538
1539        assert!(
1540            result.is_ok(),
1541            "handle_reader should exit when alive channel is closed"
1542        );
1543    }
1544
1545    /// Response-stream EOF is the normal server FIN after a closing Sentinel.
1546    /// It must not be counted as a cancellation by itself.
1547    #[tokio::test]
1548    async fn test_handle_reader_eof_does_not_count_cancellation() {
1549        let ReaderHarness {
1550            mut framed_server,
1551            framed_reader,
1552            alive_tx,
1553            alive_rx: _alive_rx,
1554            controller,
1555        } = reader_harness().await;
1556        let cancellation_counter = IntCounter::new(
1557            "tcp_client_reader_clean_eof_cancellations_test",
1558            "test cancellation counter",
1559        )
1560        .unwrap();
1561
1562        // Spawn the reader task
1563        let counter_clone = cancellation_counter.clone();
1564        let controller_clone = controller.clone();
1565        let reader_handle = tokio::spawn(async move {
1566            handle_reader(
1567                framed_reader,
1568                controller_clone,
1569                alive_tx,
1570                Some(counter_clone),
1571            )
1572            .await
1573        });
1574
1575        // Close the framed server to signal EOF to the client
1576        framed_server.close().await.unwrap();
1577
1578        // Reader should exit due to stream closure
1579        let result = tokio::time::timeout(std::time::Duration::from_secs(1), reader_handle).await;
1580
1581        assert!(
1582            result.is_ok(),
1583            "handle_reader should exit when stream is closed"
1584        );
1585        assert!(
1586            !controller.is_stopped() && !controller.is_killed(),
1587            "response-stream EOF must not cancel the context"
1588        );
1589        assert_eq!(
1590            cancellation_counter.get(),
1591            0,
1592            "response-stream EOF must not increment the cancellation counter"
1593        );
1594    }
1595
1596    /// Test that handle_reader handles multiple control messages in sequence
1597    #[tokio::test]
1598    async fn test_handle_reader_multiple_control_messages() {
1599        let ReaderHarness {
1600            mut framed_server,
1601            framed_reader,
1602            alive_tx,
1603            alive_rx: _alive_rx,
1604            controller,
1605        } = reader_harness().await;
1606
1607        // Spawn the reader task
1608        let controller_clone = controller.clone();
1609        let reader_handle = tokio::spawn(async move {
1610            handle_reader(framed_reader, controller_clone, alive_tx, None).await
1611        });
1612
1613        // Send multiple Stop messages (first one will stop, subsequent ones are no-ops)
1614        framed_server
1615            .send(control_message(&ControlMessage::Stop))
1616            .await
1617            .unwrap();
1618        framed_server
1619            .send(control_message(&ControlMessage::Stop))
1620            .await
1621            .unwrap();
1622
1623        // Close the framed server to signal EOF to the client
1624        framed_server.close().await.unwrap();
1625
1626        // Wait for reader to finish
1627        let _ = reader_handle.await.unwrap();
1628
1629        // Verify that stop was called
1630        assert!(
1631            controller.is_stopped(),
1632            "Controller should be stopped after receiving Stop messages"
1633        );
1634    }
1635
1636    /// Test handle_reader with Stop followed by Kill
1637    #[tokio::test]
1638    async fn test_handle_reader_stop_then_kill() {
1639        let ReaderHarness {
1640            mut framed_server,
1641            framed_reader,
1642            alive_tx,
1643            alive_rx: _alive_rx,
1644            controller,
1645        } = reader_harness().await;
1646
1647        // Spawn the reader task
1648        let controller_clone = controller.clone();
1649        let reader_handle = tokio::spawn(async move {
1650            handle_reader(framed_reader, controller_clone, alive_tx, None).await
1651        });
1652
1653        // Send Stop first, then Kill
1654        framed_server
1655            .send(control_message(&ControlMessage::Stop))
1656            .await
1657            .unwrap();
1658        framed_server
1659            .send(control_message(&ControlMessage::Kill))
1660            .await
1661            .unwrap();
1662
1663        // Close the framed server to signal EOF to the client
1664        framed_server.close().await.unwrap();
1665
1666        // Wait for reader to finish
1667        let _ = reader_handle.await.unwrap();
1668
1669        // Verify that kill was called (which sets killed state)
1670        assert!(
1671            controller.is_killed(),
1672            "Controller should be killed after receiving Kill message"
1673        );
1674    }
1675
1676    /// Read errors kill the context and are counted as cancellations.
1677    #[tokio::test]
1678    async fn test_handle_reader_increments_cancellation_counter_on_read_error() {
1679        let ReaderHarness {
1680            framed_server,
1681            framed_reader,
1682            alive_tx,
1683            alive_rx: _alive_rx,
1684            controller,
1685        } = reader_harness().await;
1686        let cancellation_counter = IntCounter::new(
1687            "tcp_client_reader_read_error_cancellations_test",
1688            "test cancellation counter",
1689        )
1690        .unwrap();
1691
1692        let counter_clone = cancellation_counter.clone();
1693        let controller_clone = controller.clone();
1694        let reader_handle = tokio::spawn(async move {
1695            handle_reader(
1696                framed_reader,
1697                controller_clone,
1698                alive_tx,
1699                Some(counter_clone),
1700            )
1701            .await
1702        });
1703
1704        let mut raw_writer = framed_server.into_inner();
1705        raw_writer.write_all(&[0u8; 8]).await.unwrap();
1706        raw_writer.shutdown().await.unwrap();
1707
1708        let _ = reader_handle.await.unwrap();
1709
1710        assert!(
1711            controller.is_killed(),
1712            "Controller should be killed after TCP stream read error"
1713        );
1714        assert_eq!(
1715            cancellation_counter.get(),
1716            1,
1717            "read-error close should increment cancellation metric once"
1718        );
1719    }
1720
1721    /// Drives `handle_reader` against a single message and returns the
1722    /// controller + cancellation counter for assertions.
1723    async fn run_reader_with(
1724        msg: TwoPartMessage,
1725        counter_name: &str,
1726    ) -> (Arc<Controller>, IntCounter) {
1727        let ReaderHarness {
1728            mut framed_server,
1729            framed_reader,
1730            alive_tx,
1731            alive_rx: _alive_rx,
1732            controller,
1733        } = reader_harness().await;
1734        let counter = IntCounter::new(counter_name, "test counter").unwrap();
1735
1736        let counter_clone = counter.clone();
1737        let controller_clone = controller.clone();
1738        let reader_handle = tokio::spawn(async move {
1739            handle_reader(
1740                framed_reader,
1741                controller_clone,
1742                alive_tx,
1743                Some(counter_clone),
1744            )
1745            .await
1746        });
1747
1748        framed_server.send(msg).await.unwrap();
1749        let _ = reader_handle.await.unwrap();
1750
1751        (controller, counter)
1752    }
1753
1754    /// Each protocol-violating message variant must kill only this stream
1755    /// (controller killed, cancellation counted once) and never panic the
1756    /// worker. Covers the three non-read-error panic arms in `handle_reader`:
1757    /// undecodable control bytes, server-sent Sentinel, and non-control
1758    /// (data-only) messages.
1759    #[tokio::test]
1760    async fn test_handle_reader_kills_on_protocol_violations() {
1761        let cases: Vec<(&str, TwoPartMessage)> = vec![
1762            (
1763                "invalid control bytes",
1764                TwoPartMessage::from_header(Bytes::from_static(b"not a valid control message")),
1765            ),
1766            (
1767                "sentinel from server",
1768                control_message(&ControlMessage::Sentinel),
1769            ),
1770            (
1771                "non-control (data-only)",
1772                TwoPartMessage::from_data(Bytes::from_static(b"unexpected payload")),
1773            ),
1774        ];
1775
1776        for (i, (label, msg)) in cases.into_iter().enumerate() {
1777            let counter_name = format!("tcp_client_reader_protocol_violation_test_{i}");
1778            let (controller, counter) = run_reader_with(msg, &counter_name).await;
1779            assert!(
1780                controller.is_killed(),
1781                "{label}: should kill stream context"
1782            );
1783            assert_eq!(counter.get(), 1, "{label}: should be counted once");
1784        }
1785    }
1786
1787    // ==================== handle_request_reader tests ====================
1788
1789    struct RequestReaderHarness {
1790        framed_server: FramedWrite<BoxWrite, TwoPartCodec>,
1791        framed_reader: FramedRead<BoxRead, TwoPartCodec>,
1792        bytes_tx: mpsc::Sender<Bytes>,
1793        bytes_rx: mpsc::Receiver<Bytes>,
1794        controller: Arc<Controller>,
1795    }
1796
1797    async fn request_reader_harness() -> RequestReaderHarness {
1798        let (client, server) = create_tcp_pair().await;
1799        let (read_half, _write_half) = tokio::io::split(client);
1800        let (_server_read, server_write) = tokio::io::split(server);
1801
1802        let framed_reader =
1803            FramedRead::new(Box::new(read_half) as BoxRead, TwoPartCodec::default());
1804        let framed_server =
1805            FramedWrite::new(Box::new(server_write) as BoxWrite, TwoPartCodec::default());
1806        let (bytes_tx, bytes_rx) = mpsc::channel::<Bytes>(64);
1807        let controller = Arc::new(Controller::default());
1808
1809        RequestReaderHarness {
1810            framed_server,
1811            framed_reader,
1812            bytes_tx,
1813            bytes_rx,
1814            controller,
1815        }
1816    }
1817
1818    /// Receiving Stop calls context.stop(), increments the counter, and exits.
1819    #[tokio::test]
1820    async fn test_handle_request_reader_stop_control_message() {
1821        let RequestReaderHarness {
1822            mut framed_server,
1823            framed_reader,
1824            bytes_tx,
1825            bytes_rx: _bytes_rx,
1826            controller,
1827        } = request_reader_harness().await;
1828
1829        let counter = IntCounter::new("tcp_request_reader_stop_test", "test counter").unwrap();
1830
1831        let counter_clone = counter.clone();
1832        let controller_clone = controller.clone();
1833        let handle = tokio::spawn(async move {
1834            handle_request_reader(
1835                framed_reader,
1836                bytes_tx,
1837                controller_clone,
1838                Some(counter_clone),
1839            )
1840            .await
1841        });
1842
1843        framed_server
1844            .send(control_message(&ControlMessage::Stop))
1845            .await
1846            .unwrap();
1847
1848        handle.await.unwrap();
1849
1850        assert!(controller.is_stopped(), "Stop should call context.stop()");
1851        assert!(!controller.is_killed(), "Stop should not kill the context");
1852        assert_eq!(counter.get(), 1, "cancellation counter should increment");
1853    }
1854
1855    /// Receiving Kill calls context.kill(), increments the counter, and exits.
1856    #[tokio::test]
1857    async fn test_handle_request_reader_kill_control_message() {
1858        let RequestReaderHarness {
1859            mut framed_server,
1860            framed_reader,
1861            bytes_tx,
1862            bytes_rx: _bytes_rx,
1863            controller,
1864        } = request_reader_harness().await;
1865
1866        let counter = IntCounter::new("tcp_request_reader_kill_test", "test counter").unwrap();
1867
1868        let counter_clone = counter.clone();
1869        let controller_clone = controller.clone();
1870        let handle = tokio::spawn(async move {
1871            handle_request_reader(
1872                framed_reader,
1873                bytes_tx,
1874                controller_clone,
1875                Some(counter_clone),
1876            )
1877            .await
1878        });
1879
1880        framed_server
1881            .send(control_message(&ControlMessage::Kill))
1882            .await
1883            .unwrap();
1884
1885        handle.await.unwrap();
1886
1887        assert!(controller.is_killed(), "Kill should call context.kill()");
1888        assert_eq!(counter.get(), 1, "cancellation counter should increment");
1889    }
1890
1891    /// Receiving Sentinel exits cleanly without touching the context or counter.
1892    #[tokio::test]
1893    async fn test_handle_request_reader_sentinel_control_message() {
1894        let RequestReaderHarness {
1895            mut framed_server,
1896            framed_reader,
1897            bytes_tx,
1898            mut bytes_rx,
1899            controller,
1900        } = request_reader_harness().await;
1901
1902        let counter = IntCounter::new("tcp_request_reader_sentinel_test", "test counter").unwrap();
1903
1904        let counter_clone = counter.clone();
1905        let controller_clone = controller.clone();
1906        let handle = tokio::spawn(async move {
1907            handle_request_reader(
1908                framed_reader,
1909                bytes_tx,
1910                controller_clone,
1911                Some(counter_clone),
1912            )
1913            .await
1914        });
1915
1916        framed_server
1917            .send(control_message(&ControlMessage::Sentinel))
1918            .await
1919            .unwrap();
1920
1921        handle.await.unwrap();
1922
1923        assert!(
1924            !controller.is_stopped(),
1925            "Sentinel must not stop the context"
1926        );
1927        assert!(
1928            !controller.is_killed(),
1929            "Sentinel must not kill the context"
1930        );
1931        assert_eq!(counter.get(), 0, "Sentinel must not increment counter");
1932        assert!(
1933            bytes_rx.recv().await.is_none(),
1934            "bytes_tx should be dropped on exit"
1935        );
1936    }
1937
1938    /// DataOnly frames are forwarded to bytes_tx; the loop continues until a
1939    /// terminator arrives (here, Sentinel).
1940    #[tokio::test]
1941    async fn test_handle_request_reader_forwards_data() {
1942        let RequestReaderHarness {
1943            mut framed_server,
1944            framed_reader,
1945            bytes_tx,
1946            mut bytes_rx,
1947            controller,
1948        } = request_reader_harness().await;
1949
1950        let controller_clone = controller.clone();
1951        let handle = tokio::spawn(async move {
1952            handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
1953        });
1954
1955        framed_server
1956            .send(TwoPartMessage::from_data(Bytes::from_static(b"hello")))
1957            .await
1958            .unwrap();
1959        framed_server
1960            .send(TwoPartMessage::from_data(Bytes::from_static(b"world")))
1961            .await
1962            .unwrap();
1963
1964        assert_eq!(bytes_rx.recv().await.unwrap().as_ref(), b"hello");
1965        assert_eq!(bytes_rx.recv().await.unwrap().as_ref(), b"world");
1966
1967        framed_server
1968            .send(control_message(&ControlMessage::Sentinel))
1969            .await
1970            .unwrap();
1971
1972        handle.await.unwrap();
1973        assert!(
1974            bytes_rx.recv().await.is_none(),
1975            "channel should close after Sentinel"
1976        );
1977    }
1978
1979    /// External context.kill() exits the reader without touching the wire.
1980    #[tokio::test]
1981    async fn test_handle_request_reader_exits_on_context_killed() {
1982        let RequestReaderHarness {
1983            framed_server: _framed_server,
1984            framed_reader,
1985            bytes_tx,
1986            bytes_rx: _bytes_rx,
1987            controller,
1988        } = request_reader_harness().await;
1989
1990        let controller_clone = controller.clone();
1991        let handle = tokio::spawn(async move {
1992            handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
1993        });
1994
1995        controller.kill();
1996
1997        let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
1998        assert!(
1999            result.is_ok(),
2000            "handler should exit promptly on context.kill()"
2001        );
2002    }
2003
2004    /// External context.stop() exits the reader without touching the wire.
2005    #[tokio::test]
2006    async fn test_handle_request_reader_exits_on_context_stopped() {
2007        let RequestReaderHarness {
2008            framed_server: _framed_server,
2009            framed_reader,
2010            bytes_tx,
2011            bytes_rx: _bytes_rx,
2012            controller,
2013        } = request_reader_harness().await;
2014
2015        let controller_clone = controller.clone();
2016        let handle = tokio::spawn(async move {
2017            handle_request_reader(framed_reader, bytes_tx, controller_clone, None).await
2018        });
2019
2020        controller.stop();
2021
2022        let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
2023        assert!(
2024            result.is_ok(),
2025            "handler should exit promptly on context.stop()"
2026        );
2027    }
2028
2029    /// Socket EOF exits the reader and drops bytes_tx.
2030    /// EOF before a closing Sentinel is a truncated request input: the handler
2031    /// kills the context and counts a cancellation so the consumer sees an
2032    /// aborted stream rather than a clean end.
2033    #[tokio::test]
2034    async fn test_handle_request_reader_exits_on_stream_closed() {
2035        let RequestReaderHarness {
2036            mut framed_server,
2037            framed_reader,
2038            bytes_tx,
2039            mut bytes_rx,
2040            controller,
2041        } = request_reader_harness().await;
2042
2043        let counter =
2044            IntCounter::new("tcp_request_reader_eof_truncation_test", "test counter").unwrap();
2045
2046        let counter_clone = counter.clone();
2047        let controller_clone = controller.clone();
2048        let handle = tokio::spawn(async move {
2049            handle_request_reader(
2050                framed_reader,
2051                bytes_tx,
2052                controller_clone,
2053                Some(counter_clone),
2054            )
2055            .await
2056        });
2057
2058        framed_server.close().await.unwrap();
2059
2060        let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
2061        assert!(result.is_ok(), "handler should exit on EOF");
2062        assert!(
2063            controller.is_killed(),
2064            "EOF before sentinel should kill the context (truncated input)"
2065        );
2066        assert_eq!(
2067            counter.get(),
2068            1,
2069            "EOF before sentinel should count as a cancellation"
2070        );
2071        assert!(
2072            bytes_rx.recv().await.is_none(),
2073            "bytes_tx should be dropped"
2074        );
2075    }
2076
2077    /// Dropping the returned StreamReceiver makes the reader exit promptly via
2078    /// the `bytes_tx.closed()` arm, even while parked on the socket with no
2079    /// incoming frame. This is the consumer's own choice, so it is not counted
2080    /// as a cancellation and the context is left untouched.
2081    #[tokio::test]
2082    async fn test_handle_request_reader_exits_when_receiver_dropped() {
2083        let RequestReaderHarness {
2084            framed_server,
2085            framed_reader,
2086            bytes_tx,
2087            bytes_rx,
2088            controller,
2089        } = request_reader_harness().await;
2090
2091        // Keep the socket open so the only exit path is the receiver drop.
2092        let _framed_server = framed_server;
2093
2094        let counter =
2095            IntCounter::new("tcp_request_reader_receiver_drop_test", "test counter").unwrap();
2096
2097        let counter_clone = counter.clone();
2098        let controller_clone = controller.clone();
2099        let handle = tokio::spawn(async move {
2100            handle_request_reader(
2101                framed_reader,
2102                bytes_tx,
2103                controller_clone,
2104                Some(counter_clone),
2105            )
2106            .await
2107        });
2108
2109        // Drop the consumer; the reader is parked on `framed_reader.next()`.
2110        drop(bytes_rx);
2111
2112        let result = tokio::time::timeout(std::time::Duration::from_secs(1), handle).await;
2113        assert!(
2114            result.is_ok(),
2115            "handler should exit promptly when the receiver is dropped"
2116        );
2117        assert!(
2118            !controller.is_killed() && !controller.is_stopped(),
2119            "consumer drop is not a cancellation"
2120        );
2121        assert_eq!(
2122            counter.get(),
2123            0,
2124            "consumer drop must not count as cancellation"
2125        );
2126    }
2127
2128    // ── TLS connector and SNI tests ──────────────────────────────────────────
2129
2130    fn make_ca_file() -> tempfile::NamedTempFile {
2131        use std::io::Write;
2132        let key_pair = rcgen::KeyPair::generate().unwrap();
2133        let cert = rcgen::CertificateParams::new(vec!["localhost".to_string()])
2134            .unwrap()
2135            .self_signed(&key_pair)
2136            .unwrap();
2137        let mut f = tempfile::NamedTempFile::new().unwrap();
2138        f.write_all(cert.pem().as_bytes()).unwrap();
2139        f
2140    }
2141
2142    #[test]
2143    fn connector_no_env_vars_is_plaintext() {
2144        // Clear every var the builder reads, incl. the client-identity vars, so
2145        // ambient mTLS settings can't flip this to "TLS requested".
2146        temp_env::with_vars_unset(
2147            [
2148                "DYN_TCP_TLS_CA_CERT_PATH",
2149                "DYN_TCP_TLS_INSECURE",
2150                "DYN_TCP_TLS_CLIENT_CERT_PATH",
2151                "DYN_TCP_TLS_CLIENT_KEY_PATH",
2152            ],
2153            || {
2154                assert!(build_tls_connector_from_env().unwrap().is_none());
2155            },
2156        );
2157    }
2158
2159    #[test]
2160    fn connector_insecure_is_tls() {
2161        temp_env::with_vars(
2162            [
2163                ("DYN_TCP_TLS_INSECURE", Some("true")),
2164                ("DYN_TCP_TLS_CA_CERT_PATH", None),
2165            ],
2166            || assert!(build_tls_connector_from_env().unwrap().is_some()),
2167        );
2168    }
2169
2170    #[test]
2171    fn connector_with_ca_is_tls() {
2172        let ca = make_ca_file();
2173        temp_env::with_vars(
2174            [(
2175                "DYN_TCP_TLS_CA_CERT_PATH",
2176                Some(ca.path().to_str().unwrap()),
2177            )],
2178            || assert!(build_tls_connector_from_env().unwrap().is_some()),
2179        );
2180    }
2181
2182    // Self-signed cert + matching key, for use as a client identity.
2183    fn make_identity_files() -> (tempfile::NamedTempFile, tempfile::NamedTempFile) {
2184        use std::io::Write;
2185        let key_pair = rcgen::KeyPair::generate().unwrap();
2186        let cert = rcgen::CertificateParams::new(vec!["localhost".to_string()])
2187            .unwrap()
2188            .self_signed(&key_pair)
2189            .unwrap();
2190        let mut cert_file = tempfile::NamedTempFile::new().unwrap();
2191        cert_file.write_all(cert.pem().as_bytes()).unwrap();
2192        let mut key_file = tempfile::NamedTempFile::new().unwrap();
2193        key_file
2194            .write_all(key_pair.serialize_pem().as_bytes())
2195            .unwrap();
2196        (cert_file, key_file)
2197    }
2198
2199    // CA + server leaf (SAN=localhost, serverAuth) + client leaf (clientAuth),
2200    // both signed by the CA. Returns (ca, server_cert, server_key, client_cert,
2201    // client_key) PEM temp files.
2202    #[allow(clippy::type_complexity)]
2203    fn make_mtls_chain() -> (
2204        tempfile::NamedTempFile,
2205        tempfile::NamedTempFile,
2206        tempfile::NamedTempFile,
2207        tempfile::NamedTempFile,
2208        tempfile::NamedTempFile,
2209    ) {
2210        use std::io::Write;
2211        fn write_pem(contents: &str) -> tempfile::NamedTempFile {
2212            let mut f = tempfile::NamedTempFile::new().unwrap();
2213            f.write_all(contents.as_bytes()).unwrap();
2214            f
2215        }
2216        let ca_key = rcgen::KeyPair::generate().unwrap();
2217        let mut ca_params = rcgen::CertificateParams::new(Vec::new()).unwrap();
2218        ca_params.is_ca = rcgen::IsCa::Ca(rcgen::BasicConstraints::Unconstrained);
2219        let ca_cert = ca_params.self_signed(&ca_key).unwrap();
2220
2221        let server_key = rcgen::KeyPair::generate().unwrap();
2222        let mut server_params =
2223            rcgen::CertificateParams::new(vec!["localhost".to_string()]).unwrap();
2224        server_params
2225            .extended_key_usages
2226            .push(rcgen::ExtendedKeyUsagePurpose::ServerAuth);
2227        let server_cert = server_params
2228            .signed_by(&server_key, &ca_cert, &ca_key)
2229            .unwrap();
2230
2231        let client_key = rcgen::KeyPair::generate().unwrap();
2232        let mut client_params =
2233            rcgen::CertificateParams::new(vec!["dynamo-client".to_string()]).unwrap();
2234        client_params
2235            .extended_key_usages
2236            .push(rcgen::ExtendedKeyUsagePurpose::ClientAuth);
2237        let client_cert = client_params
2238            .signed_by(&client_key, &ca_cert, &ca_key)
2239            .unwrap();
2240
2241        (
2242            write_pem(&ca_cert.pem()),
2243            write_pem(&server_cert.pem()),
2244            write_pem(&server_key.serialize_pem()),
2245            write_pem(&client_cert.pem()),
2246            write_pem(&client_key.serialize_pem()),
2247        )
2248    }
2249
2250    /// Live mTLS handshake over the **real** response-stream client path:
2251    /// `TcpClient::connect_and_split` (via an injected connector) performs the
2252    /// handshake against a server that requires a client certificate, then a
2253    /// payload round-trips over the boxed split halves. Exercises the client
2254    /// handshake + boxed I/O with mutual authentication, not a builder round-trip.
2255    #[tokio::test]
2256    async fn response_stream_client_mtls_handshake() {
2257        use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
2258
2259        let (ca, server_cert, server_key, client_cert, client_key) = make_mtls_chain();
2260
2261        // mTLS server acceptor: requires a client cert signed by the CA.
2262        let server_config = crate::tls_utils::server_tls_config(
2263            server_cert.path(),
2264            server_key.path(),
2265            Some(ca.path()),
2266        )
2267        .unwrap();
2268        let acceptor = tokio_rustls::TlsAcceptor::from(std::sync::Arc::new(server_config));
2269        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
2270        let port = listener.local_addr().unwrap().port();
2271        tokio::spawn(async move {
2272            let (tcp, _) = listener.accept().await.unwrap();
2273            let mut tls = acceptor.accept(tcp).await.expect("server mTLS handshake");
2274            let mut buf = [0u8; 5];
2275            tls.read_exact(&mut buf).await.unwrap();
2276            tls.write_all(&buf).await.unwrap();
2277            tls.flush().await.unwrap();
2278        });
2279
2280        // Real client path with an injected connector (bypasses the OnceCell);
2281        // dial by "localhost" so SNI matches the server cert's DNS SAN.
2282        let client_config = crate::tls_utils::client_tls_config(
2283            Some(ca.path()),
2284            false,
2285            Some(client_cert.path()),
2286            Some(client_key.path()),
2287        )
2288        .unwrap();
2289        let connector = tokio_rustls::TlsConnector::from(std::sync::Arc::new(client_config));
2290        let (mut reader, mut writer) = TcpClient::connect_and_split_with_connector(
2291            &format!("localhost:{port}"),
2292            Some(&connector),
2293        )
2294        .await
2295        .expect("client mTLS connect + handshake");
2296
2297        writer.write_all(b"hello").await.unwrap();
2298        writer.flush().await.unwrap();
2299        let mut got = [0u8; 5];
2300        reader.read_exact(&mut got).await.unwrap();
2301        assert_eq!(
2302            &got, b"hello",
2303            "payload should round-trip over the mutually-authenticated response-stream client path"
2304        );
2305    }
2306
2307    #[test]
2308    fn connector_partial_client_identity_errors() {
2309        // A client cert without its key must fail closed (pair validation).
2310        let ca = make_ca_file();
2311        let (client_cert, _client_key) = make_identity_files();
2312        temp_env::with_vars(
2313            [
2314                (
2315                    "DYN_TCP_TLS_CA_CERT_PATH",
2316                    Some(ca.path().to_str().unwrap()),
2317                ),
2318                (
2319                    "DYN_TCP_TLS_CLIENT_CERT_PATH",
2320                    Some(client_cert.path().to_str().unwrap()),
2321                ),
2322                ("DYN_TCP_TLS_CLIENT_KEY_PATH", None),
2323            ],
2324            || assert!(build_tls_connector_from_env().is_err()),
2325        );
2326    }
2327
2328    #[test]
2329    fn sni_parsing() {
2330        temp_env::with_var_unset("DYN_TCP_TLS_SERVER_NAME", || {
2331            assert!(matches!(
2332                tls_server_name("127.0.0.1:8080").unwrap(),
2333                ServerName::IpAddress(_)
2334            ));
2335            assert!(matches!(
2336                tls_server_name("worker-0.dynamo-system.svc.cluster.local:8080").unwrap(),
2337                ServerName::DnsName(_)
2338            ));
2339            assert!(matches!(
2340                tls_server_name("[::1]:8080").unwrap(),
2341                ServerName::IpAddress(_)
2342            ));
2343        });
2344        temp_env::with_var(
2345            "DYN_TCP_TLS_SERVER_NAME",
2346            Some("my-server.example.com"),
2347            || {
2348                assert!(matches!(
2349                    tls_server_name("127.0.0.1:8080").unwrap(),
2350                    ServerName::DnsName(_)
2351                ));
2352            },
2353        );
2354    }
2355}