1use 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}; #[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 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 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 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 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 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 framed_writer
181 .send(msg)
182 .await
183 .map_err(|e| error!("failed to send handshake: {:?}", e))?;
184
185 let (bytes_tx, bytes_rx) = tokio::sync::mpsc::channel(64);
187
188 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 tokio::spawn(async move {
202 let _ =
203 wait_for_connection_tasks(reader_task, writer_task, monitor_context, None, subject)
204 .await;
205 });
206
207 let prologue = Some(ResponseStreamPrologue {
210 error: None,
211 typed_error: None,
212 });
213
214 let stream_sender = StreamSender {
216 tx: bytes_tx,
217 prologue,
218 };
219
220 Ok(stream_sender)
221 }
222
223 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 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
295static TCP_TLS_CONNECTOR: once_cell::sync::OnceCell<Option<TlsConnector>> =
303 once_cell::sync::OnceCell::new();
304
305fn 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 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 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
364fn 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 if let Ok(sock_addr) = address.parse::<std::net::SocketAddr>() {
373 sock_addr.ip().to_string()
374 } else {
375 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 let killed = context.killed();
395 let stopped = context.stopped();
396 let bytes_closed_tx = bytes_tx.clone();
399 let bytes_closed = bytes_closed_tx.closed();
400 tokio::pin!(killed, stopped, bytes_closed);
401
402 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 _ = &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 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 cancellation_seen
499 };
500
501 if cancellation_seen && let Some(counter) = &cancellation_counter {
502 counter.inc();
503 }
504
505 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 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 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 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 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 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 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 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 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 let killed = context.killed();
697 let stopped = context.stopped();
698 tokio::pin!(killed, stopped);
699
700 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 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 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 #[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 let test_msg = TwoPartMessage::from_data(Bytes::from("test data"));
943 bytes_tx.send(test_msg).await.unwrap();
944
945 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 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 #[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 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 let mut buffer = vec![0u8; 1024];
984 let n = server.read(&mut buffer).await.unwrap();
985
986 assert!(n > 0, "Expected sentinel to be written to the TCP stream");
988
989 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 #[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 #[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 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(result);
1072
1073 let mut buffer = vec![0u8; 1024];
1075 let n = server.read(&mut buffer).await.unwrap();
1076
1077 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 #[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 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(result);
1110
1111 let mut buffer = vec![0u8; 1024];
1113 let n = server.read(&mut buffer).await.unwrap();
1114
1115 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 #[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 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 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 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 #[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 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 assert!(alive_tx.is_closed());
1185 }
1186
1187 #[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 let header_msg = TwoPartMessage::from_header(Bytes::from("header content"));
1202 bytes_tx.send(header_msg).await.unwrap();
1203
1204 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 #[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 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 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 #[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 #[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 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 #[tokio::test]
1356 async fn test_connection_monitor_aborts_writer_when_reader_panics() {
1357 let reader_task: tokio::task::JoinHandle<FramedRead<BoxRead, TwoPartCodec>> =
1361 tokio::spawn(async {
1362 panic!("simulated reader panic to trigger JoinError");
1363 });
1364
1365 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 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 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 assert!(
1402 result.unwrap().is_err(),
1403 "reader panic should propagate as Err from wait_for_connection_tasks"
1404 );
1405 }
1406
1407 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 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 #[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 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 framed_server
1463 .send(control_message(&ControlMessage::Stop))
1464 .await
1465 .unwrap();
1466
1467 framed_server.close().await.unwrap();
1469
1470 let _ = reader_handle.await.unwrap();
1472
1473 assert!(
1475 controller.is_stopped(),
1476 "Controller should be stopped after receiving Stop message"
1477 );
1478 }
1479
1480 #[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 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 framed_server
1499 .send(control_message(&ControlMessage::Kill))
1500 .await
1501 .unwrap();
1502
1503 framed_server.close().await.unwrap();
1505
1506 let _ = reader_handle.await.unwrap();
1508
1509 assert!(
1511 controller.is_killed(),
1512 "Controller should be killed after receiving Kill message"
1513 );
1514 }
1515
1516 #[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 let reader_handle =
1529 tokio::spawn(
1530 async move { handle_reader(framed_reader, controller, alive_tx, None).await },
1531 );
1532
1533 drop(alive_rx);
1535
1536 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 #[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 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 framed_server.close().await.unwrap();
1577
1578 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 #[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 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 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 framed_server.close().await.unwrap();
1625
1626 let _ = reader_handle.await.unwrap();
1628
1629 assert!(
1631 controller.is_stopped(),
1632 "Controller should be stopped after receiving Stop messages"
1633 );
1634 }
1635
1636 #[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 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 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 framed_server.close().await.unwrap();
1665
1666 let _ = reader_handle.await.unwrap();
1668
1669 assert!(
1671 controller.is_killed(),
1672 "Controller should be killed after receiving Kill message"
1673 );
1674 }
1675
1676 #[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 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 #[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 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 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(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 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 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 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 #[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 #[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 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 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 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}