1use socket2::{Domain, SockAddr, SockRef, Socket, Type};
5use std::{
6 collections::{HashMap, HashSet},
7 net::{IpAddr, SocketAddr, TcpListener},
8 os::fd::{AsFd, FromRawFd},
9 sync::Arc,
10 time::Duration,
11};
12use tokio::io::{AsyncRead, AsyncWrite};
13use tokio::time::Instant;
14use tokio_rustls::TlsAcceptor;
15
16const TOMBSTONE_TTL: Duration = Duration::from_secs(5);
21
22use bytes::Bytes;
23use derive_builder::Builder;
24use futures::{SinkExt, StreamExt};
25use local_ip_address::{Error, list_afinet_netifas, local_ip, local_ipv6};
26use parking_lot::Mutex;
27
28use serde::{Deserialize, Serialize};
29use tokio::{
30 io::AsyncWriteExt,
31 sync::{mpsc, oneshot},
32 time,
33};
34use tokio_util::codec::{FramedRead, FramedWrite};
35
36use super::{
37 CallHomeHandshake, ControlMessage, PendingConnections, RegisteredStream, StreamOptions,
38 StreamReceiver, StreamSender, TcpStreamConnectionInfo, TwoPartCodec,
39};
40use crate::discovery::EndpointInstanceId;
41use crate::engine::AsyncEngineContext;
42use crate::pipeline::{
43 PipelineError,
44 network::{
45 ResponseService, ResponseStreamPrologue, StreamPrologueError,
46 codec::{TwoPartMessage, TwoPartMessageType},
47 tcp::StreamType,
48 },
49};
50use anyhow::{Context, Result, anyhow as error};
51
52pub trait IpResolver {
54 fn local_ip(&self) -> Result<std::net::IpAddr, Error>;
55 fn local_ipv6(&self) -> Result<std::net::IpAddr, Error>;
56}
57
58pub struct DefaultIpResolver;
60
61impl IpResolver for DefaultIpResolver {
62 fn local_ip(&self) -> Result<std::net::IpAddr, Error> {
63 local_ip()
64 }
65
66 fn local_ipv6(&self) -> Result<std::net::IpAddr, Error> {
67 local_ipv6()
68 }
69}
70
71#[allow(dead_code)]
72type ResponseType = TwoPartMessage;
73
74#[derive(Debug, Serialize, Deserialize, Clone, Builder, Default)]
75pub struct ServerOptions {
76 #[builder(default = "0")]
77 pub port: u16,
78
79 #[builder(default)]
80 pub interface: Option<String>,
81}
82
83impl ServerOptions {
84 pub fn builder() -> ServerOptionsBuilder {
85 ServerOptionsBuilder::default()
86 }
87}
88
89pub struct TcpStreamServer {
93 local_ip: String,
94 local_port: u16,
95 state: Arc<Mutex<State>>,
96}
97
98#[allow(dead_code)]
105struct RequestedSendConnection {
106 context: Arc<dyn AsyncEngineContext>,
107 connection: oneshot::Sender<Result<StreamSender, StreamPrologueError>>,
108 send_buffer_count: usize,
111}
112
113struct RequestedRecvConnection {
114 context: Arc<dyn AsyncEngineContext>,
115 connection: oneshot::Sender<Result<StreamReceiver, StreamPrologueError>>,
116 send_buffer_count: usize,
119}
120
121fn data_plane_channel<T>(send_buffer_count: usize) -> (mpsc::Sender<T>, mpsc::Receiver<T>) {
127 mpsc::channel(send_buffer_count.max(1))
132}
133
134#[derive(Default)]
151struct State {
152 tx_subjects: HashMap<String, RequestedSendConnection>,
153 rx_subjects: HashMap<String, RequestedRecvConnection>,
154 subject_instance: HashMap<String, EndpointInstanceId>,
157 instance_subjects: HashMap<EndpointInstanceId, HashSet<(StreamType, String)>>,
162 removed_instances: HashMap<EndpointInstanceId, Instant>,
166 handle: Option<tokio::task::JoinHandle<Result<()>>>,
167}
168
169fn prune_tombstones(tombstones: &mut HashMap<EndpointInstanceId, Instant>, now: Instant) {
172 tombstones.retain(|_, ts| now.saturating_duration_since(*ts) < TOMBSTONE_TTL);
173}
174
175impl TcpStreamServer {
176 pub fn options_builder() -> ServerOptionsBuilder {
177 ServerOptionsBuilder::default()
178 }
179
180 pub async fn new(options: ServerOptions) -> Result<Arc<Self>, PipelineError> {
181 Self::new_with_resolver(options, DefaultIpResolver).await
182 }
183
184 pub async fn new_with_resolver<R: IpResolver>(
185 options: ServerOptions,
186 resolver: R,
187 ) -> Result<Arc<Self>, PipelineError> {
188 let local_ip = match options.interface {
189 Some(interface) => {
190 let interfaces: HashMap<String, std::net::IpAddr> =
191 list_afinet_netifas()?.into_iter().collect();
192
193 interfaces
194 .get(&interface)
195 .ok_or(PipelineError::Generic(format!(
196 "Interface not found: {}",
197 interface
198 )))?
199 .to_string()
200 }
201 None => {
202 let resolved_ip = resolver.local_ip().or_else(|err| match err {
203 Error::LocalIpAddressNotFound => resolver.local_ipv6(),
204 _ => Err(err),
205 });
206
207 match resolved_ip {
208 Ok(addr) => addr,
209 Err(Error::LocalIpAddressNotFound) => {
214 tracing::warn!(
215 "No routable local IP address found; falling back to 127.0.0.1"
216 );
217 IpAddr::from([127, 0, 0, 1])
218 }
219 Err(err) => {
220 return Err(PipelineError::Generic(format!(
221 "Failed to resolve local IP address: {err}"
222 )));
223 }
224 }
225 .to_string()
226 }
227 };
228
229 let state = Arc::new(Mutex::new(State::default()));
230
231 let tls_acceptor = Self::build_tls_acceptor().map_err(|e| {
233 PipelineError::Generic(format!("Failed to build TCP TLS acceptor: {}", e))
234 })?;
235
236 let local_port = Self::start(local_ip.clone(), options.port, state.clone(), tls_acceptor)
237 .await
238 .map_err(|e| {
239 PipelineError::Generic(format!("Failed to start TcpStreamServer: {}", e))
240 })?;
241
242 tracing::debug!("tcp transport service on {local_ip}:{local_port}");
243
244 Ok(Arc::new(Self {
245 local_ip,
246 local_port,
247 state,
248 }))
249 }
250
251 pub async fn associate_instance(
264 &self,
265 recv_subject: &str,
266 send_subject: Option<&str>,
267 id: &EndpointInstanceId,
268 ) -> bool {
269 let mut state = self.state.lock();
270 let now = Instant::now();
271 prune_tombstones(&mut state.removed_instances, now);
272 if state.removed_instances.contains_key(id) {
273 tracing::warn!(
275 recv_subject,
276 send_subject,
277 namespace = %id.namespace,
278 component = %id.component,
279 endpoint = %id.endpoint,
280 instance_id = id.instance_id,
281 "Cancelling subject immediately: instance already removed (tombstoned)"
282 );
283 state.rx_subjects.remove(recv_subject);
284 if let Some(s) = send_subject {
285 state.tx_subjects.remove(s);
286 }
287 return false;
288 }
289 state
290 .subject_instance
291 .insert(recv_subject.to_string(), id.clone());
292 if let Some(s) = send_subject {
293 state.subject_instance.insert(s.to_string(), id.clone());
294 }
295 let entry = state.instance_subjects.entry(id.clone()).or_default();
296 entry.insert((StreamType::Response, recv_subject.to_string()));
297 if let Some(s) = send_subject {
298 entry.insert((StreamType::Request, s.to_string()));
299 }
300 true
301 }
302
303 pub async fn cancel_recv_stream(&self, subject: &str) {
306 let mut state = self.state.lock();
307 state.rx_subjects.remove(subject);
308 if let Some(key) = state.subject_instance.remove(subject)
309 && let Some(subjects) = state.instance_subjects.get_mut(&key)
310 {
311 subjects.remove(&(StreamType::Response, subject.to_string()));
312 if subjects.is_empty() {
313 state.instance_subjects.remove(&key);
314 }
315 }
316 }
317
318 pub async fn cancel_send_stream(&self, subject: &str) {
324 let mut state = self.state.lock();
325 state.tx_subjects.remove(subject);
326 if let Some(key) = state.subject_instance.remove(subject)
327 && let Some(subjects) = state.instance_subjects.get_mut(&key)
328 {
329 subjects.remove(&(StreamType::Request, subject.to_string()));
330 if subjects.is_empty() {
331 state.instance_subjects.remove(&key);
332 }
333 }
334 }
335
336 pub async fn cancel_instance_streams(&self, id: &EndpointInstanceId) -> usize {
341 let mut state = self.state.lock();
342 let now = Instant::now();
343 prune_tombstones(&mut state.removed_instances, now);
344 state.removed_instances.insert(id.clone(), now);
345 let subjects = match state.instance_subjects.remove(id) {
346 Some(subjects) => subjects,
347 None => return 0,
348 };
349 let count = subjects.len();
350 for (kind, subject) in &subjects {
351 match kind {
352 StreamType::Response => {
353 state.rx_subjects.remove(subject);
354 }
355 StreamType::Request => {
356 state.tx_subjects.remove(subject);
357 }
358 }
359 state.subject_instance.remove(subject);
360 }
361 count
362 }
363
364 pub async fn clear_instance_tombstone(&self, id: &EndpointInstanceId) {
367 let mut state = self.state.lock();
368 state.removed_instances.remove(id);
369 }
370
371 fn build_tls_acceptor() -> anyhow::Result<Option<TlsAcceptor>> {
376 use crate::config::environment_names::tcp_response_stream::tls as env;
377 let cert_path = std::env::var(env::DYN_TCP_TLS_CERT_PATH).ok();
378 let key_path = std::env::var(env::DYN_TCP_TLS_KEY_PATH).ok();
379 let client_ca = std::env::var(env::DYN_TCP_TLS_CLIENT_CA_CERT_PATH).ok();
380 Ok(crate::tls_utils::server_tls_acceptor_config(
381 "TCP server",
382 cert_path.as_deref().map(std::path::Path::new),
383 key_path.as_deref().map(std::path::Path::new),
384 client_ca.as_deref().map(std::path::Path::new),
385 )?
386 .map(|config| TlsAcceptor::from(Arc::new(config))))
387 }
388
389 async fn start(
390 local_ip: String,
391 local_port: u16,
392 state: Arc<Mutex<State>>,
393 tls_acceptor: Option<TlsAcceptor>,
394 ) -> Result<u16> {
395 let addr = format!("{}:{}", local_ip, local_port);
396 let state_clone = state.clone();
397 let (ready_tx, ready_rx) = tokio::sync::oneshot::channel::<Result<u16>>();
398 {
399 let mut guard = state.lock();
400 if guard.handle.is_some() {
401 panic!("TcpStreamServer already started");
402 }
403 guard.handle = Some(tokio::spawn(tcp_listener(
404 addr,
405 state_clone,
406 tls_acceptor,
407 ready_tx,
408 )));
409 }
410 let local_port = ready_rx.await??;
411 Ok(local_port)
412 }
413
414 fn insert_request_stream(&self, subject: String, connection: RequestedSendConnection) {
415 self.state.lock().tx_subjects.insert(subject, connection);
416 }
417
418 fn insert_response_stream(&self, subject: String, connection: RequestedRecvConnection) {
419 self.state.lock().rx_subjects.insert(subject, connection);
420 }
421
422 fn take_request_stream(state: &Mutex<State>, subject: &str) -> Option<RequestedSendConnection> {
423 let mut state = state.lock();
424 let connection = state.tx_subjects.remove(subject);
425 if let Some(key) = state.subject_instance.remove(subject)
426 && let Some(subjects) = state.instance_subjects.get_mut(&key)
427 {
428 subjects.remove(&(StreamType::Request, subject.to_string()));
429 if subjects.is_empty() {
430 state.instance_subjects.remove(&key);
431 }
432 }
433 connection
434 }
435
436 fn take_response_stream(
437 state: &Mutex<State>,
438 subject: &str,
439 ) -> Option<RequestedRecvConnection> {
440 let mut state = state.lock();
441 let connection = state.rx_subjects.remove(subject);
442 if let Some(key) = state.subject_instance.remove(subject)
443 && let Some(subjects) = state.instance_subjects.get_mut(&key)
444 {
445 subjects.remove(&(StreamType::Response, subject.to_string()));
446 if subjects.is_empty() {
447 state.instance_subjects.remove(&key);
448 }
449 }
450 connection
451 }
452}
453
454#[async_trait::async_trait]
456impl ResponseService for TcpStreamServer {
457 async fn register(&self, options: StreamOptions) -> PendingConnections {
478 let address = format!("{}:{}", self.local_ip, self.local_port);
481 tracing::debug!("Registering new TcpStream on {address}");
482
483 let send_stream = if options.enable_request_stream {
484 let sender_subject = uuid::Uuid::new_v4().to_string();
485 let registry_subject = sender_subject.clone();
486
487 let (pending_sender_tx, pending_sender_rx) = oneshot::channel();
488
489 let connection_info = RequestedSendConnection {
490 context: options.context.clone(),
491 connection: pending_sender_tx,
492 send_buffer_count: options.send_buffer_count,
493 };
494
495 let cleanup_subject = sender_subject.clone();
496 let cleanup_state = self.state.clone();
497 let registered_stream = RegisteredStream::new(
498 TcpStreamConnectionInfo {
499 address: address.clone(),
500 subject: sender_subject,
501 context: options.context.id().to_string(),
502 stream_type: StreamType::Request,
503 }
504 .into(),
505 pending_sender_rx,
506 )
507 .with_cleanup(move || {
508 tokio::spawn(async move {
510 let mut state = cleanup_state.lock();
511 state.tx_subjects.remove(&cleanup_subject);
512 if let Some(key) = state.subject_instance.remove(&cleanup_subject)
513 && let Some(subjects) = state.instance_subjects.get_mut(&key)
514 {
515 subjects.remove(&(StreamType::Request, cleanup_subject.clone()));
516 if subjects.is_empty() {
517 state.instance_subjects.remove(&key);
518 }
519 }
520 });
521 });
522
523 self.insert_request_stream(registry_subject, connection_info);
524
525 Some(registered_stream)
526 } else {
527 None
528 };
529
530 let recv_stream = if options.enable_response_stream {
531 let (pending_recver_tx, pending_recver_rx) = oneshot::channel();
532 let receiver_subject = uuid::Uuid::new_v4().to_string();
533 let registry_subject = receiver_subject.clone();
534
535 let connection_info = RequestedRecvConnection {
536 context: options.context.clone(),
537 connection: pending_recver_tx,
538 send_buffer_count: options.send_buffer_count,
539 };
540
541 let cleanup_subject = receiver_subject.clone();
542 let cleanup_state = self.state.clone();
543 let registered_stream = RegisteredStream::new(
544 TcpStreamConnectionInfo {
545 address: address.clone(),
546 subject: receiver_subject,
547 context: options.context.id().to_string(),
548 stream_type: StreamType::Response,
549 }
550 .into(),
551 pending_recver_rx,
552 )
553 .with_cleanup(move || {
554 tokio::spawn(async move {
556 let mut state = cleanup_state.lock();
557 state.rx_subjects.remove(&cleanup_subject);
558 if let Some(key) = state.subject_instance.remove(&cleanup_subject)
559 && let Some(subjects) = state.instance_subjects.get_mut(&key)
560 {
561 subjects.remove(&(StreamType::Response, cleanup_subject.clone()));
562 if subjects.is_empty() {
563 state.instance_subjects.remove(&key);
564 }
565 }
566 });
567 });
568
569 self.insert_response_stream(registry_subject, connection_info);
570
571 Some(registered_stream)
572 } else {
573 None
574 };
575
576 PendingConnections {
577 send_stream,
578 recv_stream,
579 }
580 }
581}
582
583const ACCEPT_BACKOFF_INITIAL_DELAY: Duration = Duration::from_millis(5);
585const ACCEPT_BACKOFF_MAX_DELAY: Duration = Duration::from_secs(1);
588const ACCEPT_BACKOFF_LOG_INTERVAL: Duration = Duration::from_secs(5);
590
591#[derive(Debug, Clone, Copy, PartialEq, Eq)]
593enum AcceptFailure {
594 Exhaustion,
598 Ordinary,
601}
602
603#[derive(Debug, Clone, Copy, PartialEq, Eq)]
605struct AcceptBackoffAction {
606 delay: Duration,
608 log_suppressed: Option<u64>,
612}
613
614#[derive(Debug)]
621struct AcceptBackoff {
622 initial_delay: Duration,
623 max_delay: Duration,
624 log_interval: Duration,
625 current_delay: Duration,
626 suppressed: u64,
627 last_log_at: Option<std::time::Instant>,
628 in_backoff: bool,
632}
633
634impl Default for AcceptBackoff {
635 fn default() -> Self {
636 Self {
637 initial_delay: ACCEPT_BACKOFF_INITIAL_DELAY,
638 max_delay: ACCEPT_BACKOFF_MAX_DELAY,
639 log_interval: ACCEPT_BACKOFF_LOG_INTERVAL,
640 current_delay: ACCEPT_BACKOFF_INITIAL_DELAY,
641 suppressed: 0,
642 last_log_at: None,
643 in_backoff: false,
644 }
645 }
646}
647
648impl AcceptBackoff {
649 fn classify(err: &std::io::Error) -> AcceptFailure {
653 #[cfg(unix)]
654 {
655 if matches!(
661 err.raw_os_error(),
662 Some(libc::EMFILE) | Some(libc::ENFILE) | Some(libc::ENOBUFS) | Some(libc::ENOMEM)
663 ) {
664 return AcceptFailure::Exhaustion;
665 }
666 }
667 #[cfg(not(unix))]
668 {
669 let _ = err;
670 }
671 AcceptFailure::Ordinary
672 }
673
674 fn record_exhaustion(&mut self, now: std::time::Instant) -> AcceptBackoffAction {
678 self.in_backoff = true;
679
680 let delay = self.current_delay;
681 self.current_delay = self.current_delay.saturating_mul(2).min(self.max_delay);
682
683 let due = match self.last_log_at {
684 None => true,
685 Some(prev) => now.saturating_duration_since(prev) >= self.log_interval,
686 };
687
688 let log_suppressed = if due {
689 self.last_log_at = Some(now);
690 Some(std::mem::take(&mut self.suppressed))
691 } else {
692 self.suppressed += 1;
693 None
694 };
695
696 AcceptBackoffAction {
697 delay,
698 log_suppressed,
699 }
700 }
701
702 fn record_success(&mut self, now: impl FnOnce() -> std::time::Instant) -> Option<u64> {
711 if !self.in_backoff {
712 return None;
713 }
714 self.current_delay = self.initial_delay;
715 self.in_backoff = false;
716 let now = now();
717
718 let due = match self.last_log_at {
719 None => true,
720 Some(prev) => now.saturating_duration_since(prev) >= self.log_interval,
721 };
722 if !due {
723 return None;
724 }
725 self.last_log_at = Some(now);
726 Some(std::mem::take(&mut self.suppressed))
727 }
728}
729
730async fn handle_accept_error(err: &std::io::Error, backoff: &mut AcceptBackoff) -> Duration {
734 match AcceptBackoff::classify(err) {
735 AcceptFailure::Ordinary => {
736 tracing::warn!(error = %err, "failed to accept tcp connection");
738 #[cfg(debug_assertions)]
741 eprintln!("failed to accept tcp connection: {}", err);
742 Duration::ZERO
743 }
744 AcceptFailure::Exhaustion => {
745 crate::metrics::transport_metrics::TCP_ACCEPT_BACKOFF_TOTAL.inc();
746 let action = backoff.record_exhaustion(std::time::Instant::now());
747 if let Some(suppressed) = action.log_suppressed {
748 tracing::warn!(
749 error = %err,
750 retry_delay_ms = action.delay.as_millis() as u64,
751 suppressed_failures = suppressed,
752 "tcp accept failed: out of file descriptors or kernel memory; backing off before retry"
753 );
754 }
755 time::sleep(action.delay).await;
756 action.delay
757 }
758 }
759}
760
761type BoxRead = Box<dyn tokio::io::AsyncRead + Unpin + Send>;
763type BoxWrite = Box<dyn tokio::io::AsyncWrite + Unpin + Send>;
764
765async fn tcp_listener(
772 addr: String,
773 state: Arc<Mutex<State>>,
774 tls_acceptor: Option<TlsAcceptor>,
775 read_tx: tokio::sync::oneshot::Sender<Result<u16>>,
776) -> Result<()> {
777 let listener = tokio::net::TcpListener::bind(&addr)
778 .await
779 .map_err(|e| anyhow::anyhow!("Failed to start TcpListender on {}: {}", addr, e));
780
781 let listener = match listener {
782 Ok(listener) => {
783 let addr = listener
784 .local_addr()
785 .map_err(|e| anyhow::anyhow!("Failed get SocketAddr: {:?}", e))
786 .unwrap();
787
788 read_tx
789 .send(Ok(addr.port()))
790 .expect("Failed to send ready signal");
791
792 listener
793 }
794 Err(e) => {
795 read_tx.send(Err(e)).expect("Failed to send ready signal");
796 return Err(anyhow::anyhow!("Failed to start TcpListender on {}", addr));
797 }
798 };
799
800 let mut accept_backoff = AcceptBackoff::default();
801
802 loop {
803 let (stream, _addr) = match listener.accept().await {
809 Ok((stream, _addr)) => {
810 if let Some(suppressed) = accept_backoff.record_success(std::time::Instant::now) {
811 tracing::warn!(
812 suppressed_failures = suppressed,
813 "tcp accept recovered from resource exhaustion"
814 );
815 }
816 (stream, _addr)
817 }
818 Err(e) => {
819 handle_accept_error(&e, &mut accept_backoff).await;
820 continue;
821 }
822 };
823
824 match stream.set_nodelay(true) {
825 Ok(_) => (),
826 Err(e) => {
827 tracing::warn!("failed to set tcp stream to nodelay: {e}");
828 }
829 }
830
831 match SockRef::from(&stream).set_linger(Some(std::time::Duration::from_secs(0))) {
832 Ok(_) => (),
833 Err(e) => {
834 tracing::warn!("failed to set tcp stream to linger: {e}");
835 }
836 }
837
838 let state_clone = state.clone();
841 let tls_acceptor_clone = tls_acceptor.clone();
842 tokio::spawn(async move {
843 let (reader, writer) = if let Some(ref tls) = tls_acceptor_clone {
844 match tokio::time::timeout(
845 crate::tls_utils::handshake_timeout(),
846 tls.accept(stream),
847 )
848 .await
849 {
850 Ok(Ok(tls_stream)) => {
851 let (r, w) = tokio::io::split(tls_stream);
852 (Box::new(r) as BoxRead, Box::new(w) as BoxWrite)
853 }
854 Ok(Err(e)) => {
855 tracing::warn!("TLS handshake failed: {e}");
856 return;
857 }
858 Err(_) => {
859 tracing::warn!("TLS handshake timed out");
860 return;
861 }
862 }
863 } else {
864 let (r, w) = tokio::io::split(stream);
865 (Box::new(r) as BoxRead, Box::new(w) as BoxWrite)
866 };
867 handle_connection(reader, writer, state_clone).await;
868 });
869 }
870
871 async fn handle_connection(reader: BoxRead, writer: BoxWrite, state: Arc<Mutex<State>>) {
874 let result = process_stream(reader, writer, state).await;
875 match result {
876 Ok(_) => tracing::trace!("successfully processed tcp connection"),
877 Err(e) => {
878 tracing::warn!("failed to handle tcp connection: {e}");
879 #[cfg(debug_assertions)]
880 eprintln!("failed to handle tcp connection: {}", e);
881 }
882 }
883 }
884
885 async fn process_stream(
888 read_half: BoxRead,
889 write_half: BoxWrite,
890 state: Arc<Mutex<State>>,
891 ) -> Result<()> {
892 let mut framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
894 let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
895
896 let first_message =
899 tokio::time::timeout(std::time::Duration::from_secs(10), framed_reader.next())
900 .await
901 .map_err(|_| error!("Timed out waiting for CallHomeHandshake"))?
902 .ok_or(error!("Connection closed without a ControlMessage"))??;
903
904 let handshake: CallHomeHandshake = match first_message.header() {
907 Some(header) => serde_json::from_slice(header).map_err(|e| {
908 error!(
909 "Failed to deserialize the first message as a valid `CallHomeHandshake`: {e}",
910 )
911 })?,
912 None => {
913 return Err(error!("Expected ControlMessage, got DataMessage"));
914 }
915 };
916
917 match handshake.stream_type {
919 StreamType::Request => {
920 process_request_stream(handshake.subject, state, framed_reader, framed_writer).await
921 }
922 StreamType::Response => {
923 process_response_stream(handshake.subject, state, framed_reader, framed_writer)
924 .await
925 }
926 }
927 }
928
929 async fn process_request_stream(
940 subject: String,
941 state: Arc<Mutex<State>>,
942 reader: FramedRead<BoxRead, TwoPartCodec>,
943 writer: FramedWrite<BoxWrite, TwoPartCodec>,
944 ) -> Result<()> {
945 drop(reader);
947
948 let request_stream = TcpStreamServer::take_request_stream(&state, &subject).ok_or_else(|| {
949 error!(
950 "Subject not found: {}; downstream subscriber specified a subject unknown to the upstream publisher",
951 subject
952 )
953 })?;
954
955 let RequestedSendConnection {
956 context,
957 connection,
958 send_buffer_count,
959 } = request_stream;
960
961 let (request_tx, request_rx) = data_plane_channel(send_buffer_count);
965
966 if connection
967 .send(Ok(crate::pipeline::network::StreamSender {
968 tx: request_tx,
969 prologue: None,
972 }))
973 .is_err()
974 {
975 return Err(error!(
976 "The requester of the request stream has been dropped before the connection was established"
977 ));
978 }
979
980 request_stream_send_handler(writer, request_rx, context).await;
981 Ok(())
982 }
983
984 async fn request_stream_send_handler(
994 mut framed_writer: FramedWrite<BoxWrite, TwoPartCodec>,
995 mut request_rx: mpsc::Receiver<TwoPartMessage>,
996 context: Arc<dyn AsyncEngineContext>,
997 ) {
998 let killed = context.killed();
1001 let stopped = context.stopped();
1002 tokio::pin!(killed, stopped);
1003
1004 let closing_msg: Option<ControlMessage> = loop {
1005 tokio::select! {
1006 biased;
1007
1008 _ = &mut killed => {
1009 tracing::trace!("context kill received in request-stream send handler");
1010 break Some(ControlMessage::Kill);
1011 }
1012
1013 _ = &mut stopped => {
1014 tracing::trace!("context stop received in request-stream send handler");
1015 break Some(ControlMessage::Stop);
1016 }
1017
1018 msg = request_rx.recv() => {
1019 match msg {
1020 Some(msg) => {
1021 if let Err(e) = framed_writer.send(msg).await {
1022 tracing::trace!(
1023 "failed to send request-stream frame to downstream: {:?}",
1024 e
1025 );
1026 break None;
1027 }
1028 }
1029 None => {
1030 tracing::trace!("upstream request-stream sender closed; sending sentinel");
1031 break Some(ControlMessage::Sentinel);
1032 }
1033 }
1034 }
1035 }
1036 };
1037
1038 if let Some(ctrl) = closing_msg
1039 && let Ok(bytes) = serde_json::to_vec(&ctrl)
1040 && let Err(err) = framed_writer
1041 .send(TwoPartMessage::from_header(bytes.into()))
1042 .await
1043 {
1044 tracing::trace!(?err, ?ctrl, "request-stream closing-frame send failed");
1045 }
1046
1047 let mut inner = framed_writer.into_inner();
1048 if let Err(err) = inner.flush().await {
1049 tracing::trace!(?err, "request-stream socket flush failed");
1050 }
1051 if let Err(err) = inner.shutdown().await {
1052 tracing::trace!(?err, "request-stream socket shutdown failed");
1053 }
1054 }
1055
1056 async fn process_response_stream(
1057 subject: String,
1058 state: Arc<Mutex<State>>,
1059 mut reader: FramedRead<BoxRead, TwoPartCodec>,
1060 writer: FramedWrite<BoxWrite, TwoPartCodec>,
1061 ) -> Result<()> {
1062 let response_stream = TcpStreamServer::take_response_stream(&state, &subject).ok_or_else(|| {
1063 error!("Subject not found: {}; upstream publisher specified a subject unknown to the downsteam subscriber", subject)
1064 })?;
1065
1066 let RequestedRecvConnection {
1068 context,
1069 connection,
1070 send_buffer_count,
1071 } = response_stream;
1072
1073 let prologue = reader
1078 .next()
1079 .await
1080 .ok_or(error!("Connection closed without a ControlMessge"))??;
1081
1082 let prologue = match prologue.into_message_type() {
1084 TwoPartMessageType::HeaderOnly(header) => {
1085 match serde_json::from_slice::<ResponseStreamPrologue>(&header) {
1086 Ok(prologue) => prologue,
1087 Err(e) => {
1088 let msg = format!("malformed prologue: {e}");
1093 let _ =
1094 connection.send(Err(StreamPrologueError::from_message(msg.clone())));
1095 return Err(error!(msg));
1096 }
1097 }
1098 }
1099 _ => {
1100 let msg = "malformed prologue: expected HeaderOnly ControlMessage";
1105 let _ = connection.send(Err(StreamPrologueError::from_message(msg)));
1106 return Err(error!(msg));
1107 }
1108 };
1109
1110 if let Some(error) = prologue.error {
1117 let returned = error!("Received error prologue: {error}");
1118 let _ = connection.send(Err(StreamPrologueError {
1121 message: error,
1122 typed_error: prologue.typed_error,
1123 }));
1124 return Err(returned);
1125 }
1126
1127 let (response_tx, response_rx) = data_plane_channel(send_buffer_count);
1131
1132 if connection
1133 .send(Ok(crate::pipeline::network::StreamReceiver {
1134 rx: response_rx,
1135 }))
1136 .is_err()
1137 {
1138 return Err(error!(
1139 "The requester of the stream has been dropped before the connection was established"
1140 ));
1141 }
1142
1143 let (control_tx, control_rx) = mpsc::channel::<ControlMessage>(1);
1144
1145 let send_task = tokio::spawn(network_send_handler(writer, control_rx));
1149
1150 let recv_task = tokio::spawn(network_receive_handler(
1152 reader,
1153 response_tx,
1154 control_tx,
1155 context.clone(),
1156 ));
1157
1158 let (monitor_result, forward_result) = tokio::join!(send_task, recv_task);
1160
1161 monitor_result?;
1162 forward_result?;
1163
1164 Ok(())
1165 }
1166
1167 async fn network_receive_handler(
1168 mut framed_reader: FramedRead<BoxRead, TwoPartCodec>,
1169 response_tx: mpsc::Sender<Bytes>,
1170 control_tx: mpsc::Sender<ControlMessage>,
1171 context: Arc<dyn AsyncEngineContext>,
1172 ) {
1173 let response_closed = response_tx.closed();
1176 let killed = context.killed();
1177 let stopped = context.stopped();
1178 tokio::pin!(response_closed, killed, stopped);
1179
1180 let mut can_stop = true;
1182 loop {
1183 tokio::select! {
1184 biased;
1185
1186 _ = &mut response_closed => {
1187 tracing::trace!("response channel closed before the client finished writing data");
1188 let _ = control_tx.send(ControlMessage::Kill).await;
1189 break;
1190 }
1191
1192 _ = &mut killed => {
1193 tracing::trace!("context kill signal received; shutting down");
1194 let _ = control_tx.send(ControlMessage::Kill).await;
1195 break;
1196 }
1197
1198 _ = &mut stopped, if can_stop => {
1199 tracing::trace!("context stop signal received; shutting down");
1200 can_stop = false;
1203 let _ = control_tx.send(ControlMessage::Stop).await;
1204 }
1205
1206 msg = framed_reader.next() => {
1207 match msg {
1208 Some(Ok(msg)) => {
1209 let (header, data) = msg.into_parts();
1210
1211 if !header.is_empty() {
1213 match process_control_message(header) {
1214 Ok(ControlAction::Continue) => {}
1215 Ok(ControlAction::Shutdown) => {
1216 if !data.is_empty() {
1217 tracing::warn!(
1221 data_len = data.len(),
1222 "client sent Sentinel with data (protocol violation); killing stream"
1223 );
1224 let _ = control_tx.send(ControlMessage::Kill).await;
1225 break;
1226 }
1227 tracing::trace!("received sentinel message; shutting down");
1228 break;
1229 }
1230 Err(e) => {
1231 tracing::warn!(err = ?e, "malformed control message, closing connection");
1234 let _ = control_tx.send(ControlMessage::Kill).await;
1235 break;
1236 }
1237 }
1238 }
1239
1240 if !data.is_empty()
1241 && let Err(err) = response_tx.send(data).await {
1242 tracing::debug!(?err, "forwarding body/data to response channel failed");
1243 let _ = control_tx.send(ControlMessage::Kill).await;
1244 break;
1245 };
1246 }
1247 Some(Err(e)) => {
1248 tracing::warn!(err = ?e, "tcp stream read error from worker, closing connection");
1251 let _ = control_tx.send(ControlMessage::Kill).await;
1252 break;
1253 }
1254 None => {
1255 tracing::trace!("tcp stream was closed by client");
1261 break;
1262 }
1263 }
1264 }
1265
1266 }
1267 }
1268 }
1269
1270 async fn network_send_handler(
1271 socket_tx: FramedWrite<BoxWrite, TwoPartCodec>,
1272 control_rx: mpsc::Receiver<ControlMessage>,
1273 ) {
1274 let mut socket_tx = socket_tx;
1275 let mut control_rx = control_rx;
1276
1277 while let Some(control_msg) = control_rx.recv().await {
1278 if matches!(control_msg, ControlMessage::Sentinel) {
1282 tracing::warn!("received sentinel on send-side control channel; dropping");
1283 continue;
1284 }
1285 let bytes = match serde_json::to_vec(&control_msg) {
1286 Ok(b) => b,
1287 Err(e) => {
1288 tracing::warn!(err = ?e, ?control_msg, "failed to serialize control message");
1291 continue;
1292 }
1293 };
1294 let message = TwoPartMessage::from_header(bytes.into());
1295 match socket_tx.send(message).await {
1296 Ok(_) => tracing::debug!(?control_msg, "issued control message"),
1297 Err(e) => {
1298 tracing::debug!(err = ?e, ?control_msg, "failed to send control message")
1299 }
1300 }
1301 }
1302
1303 let mut inner = socket_tx.into_inner();
1304 if let Err(e) = inner.flush().await {
1305 tracing::debug!("failed to flush socket: {e}");
1306 }
1307 if let Err(e) = inner.shutdown().await {
1308 tracing::debug!("failed to shutdown socket: {e}");
1309 }
1310 }
1311}
1312
1313enum ControlAction {
1314 Continue,
1315 Shutdown,
1316}
1317
1318fn process_control_message(message: Bytes) -> Result<ControlAction> {
1319 match serde_json::from_slice::<ControlMessage>(&message)? {
1320 ControlMessage::Sentinel => {
1321 tracing::trace!("sentinel received; shutting down");
1324 Ok(ControlAction::Shutdown)
1325 }
1326 ControlMessage::Kill | ControlMessage::Stop => {
1327 anyhow::bail!("unexpected control message on response stream");
1331 }
1332 }
1333}
1334
1335#[cfg(test)]
1336mod tests {
1337 use super::*;
1338 use crate::engine::AsyncEngineContextProvider;
1339 use crate::error::{BackendError, DynamoError, ErrorType};
1340 use crate::pipeline::Context;
1341 use crate::pipeline::network::DEFAULT_SEND_BUFFER_COUNT;
1342 use crate::pipeline::network::tcp::client::TcpClient;
1343 use std::io::Write;
1344 use tempfile::NamedTempFile;
1345 use tokio::io::{AsyncWriteExt, ReadHalf, WriteHalf};
1346 use tokio::net::TcpStream;
1347
1348 fn make_cert_files() -> (NamedTempFile, NamedTempFile) {
1349 let key_pair = rcgen::KeyPair::generate().unwrap();
1350 let cert = rcgen::CertificateParams::new(vec!["localhost".to_string()])
1351 .unwrap()
1352 .self_signed(&key_pair)
1353 .unwrap();
1354 let mut cert_file = NamedTempFile::new().unwrap();
1355 cert_file.write_all(cert.pem().as_bytes()).unwrap();
1356 let mut key_file = NamedTempFile::new().unwrap();
1357 key_file
1358 .write_all(key_pair.serialize_pem().as_bytes())
1359 .unwrap();
1360 (cert_file, key_file)
1361 }
1362
1363 #[test]
1364 fn build_tls_acceptor_no_env_vars_is_plaintext() {
1365 temp_env::with_vars_unset(
1368 [
1369 "DYN_TCP_TLS_CERT_PATH",
1370 "DYN_TCP_TLS_KEY_PATH",
1371 "DYN_TCP_TLS_CLIENT_CA_CERT_PATH",
1372 ],
1373 || {
1374 assert!(TcpStreamServer::build_tls_acceptor().unwrap().is_none());
1375 },
1376 );
1377 }
1378
1379 #[test]
1380 fn build_tls_acceptor_partial_config_errors() {
1381 let (cert, key) = make_cert_files();
1382 let cert_str = cert.path().to_str().unwrap();
1383 let key_str = key.path().to_str().unwrap();
1384 temp_env::with_vars(
1386 [
1387 ("DYN_TCP_TLS_CERT_PATH", Some(cert_str)),
1388 ("DYN_TCP_TLS_KEY_PATH", None),
1389 ],
1390 || assert!(TcpStreamServer::build_tls_acceptor().is_err()),
1391 );
1392 temp_env::with_vars(
1394 [
1395 ("DYN_TCP_TLS_CERT_PATH", None),
1396 ("DYN_TCP_TLS_KEY_PATH", Some(key_str)),
1397 ],
1398 || assert!(TcpStreamServer::build_tls_acceptor().is_err()),
1399 );
1400 }
1401
1402 #[test]
1403 fn build_tls_acceptor_both_paths_is_tls() {
1404 let (cert, key) = make_cert_files();
1405 temp_env::with_vars(
1406 [
1407 ("DYN_TCP_TLS_CERT_PATH", Some(cert.path().to_str().unwrap())),
1408 ("DYN_TCP_TLS_KEY_PATH", Some(key.path().to_str().unwrap())),
1409 ],
1410 || assert!(TcpStreamServer::build_tls_acceptor().unwrap().is_some()),
1411 );
1412 }
1413
1414 #[test]
1415 fn build_tls_acceptor_with_client_ca_is_mtls() {
1416 let (cert, key) = make_cert_files();
1418 temp_env::with_vars(
1419 [
1420 ("DYN_TCP_TLS_CERT_PATH", Some(cert.path().to_str().unwrap())),
1421 ("DYN_TCP_TLS_KEY_PATH", Some(key.path().to_str().unwrap())),
1422 (
1423 "DYN_TCP_TLS_CLIENT_CA_CERT_PATH",
1424 Some(cert.path().to_str().unwrap()),
1425 ),
1426 ],
1427 || assert!(TcpStreamServer::build_tls_acceptor().unwrap().is_some()),
1428 );
1429 }
1430
1431 #[test]
1432 fn build_tls_acceptor_client_ca_without_server_identity_errors() {
1433 let (cert, _key) = make_cert_files();
1434 temp_env::with_vars(
1435 [
1436 ("DYN_TCP_TLS_CERT_PATH", None),
1437 ("DYN_TCP_TLS_KEY_PATH", None),
1438 (
1439 "DYN_TCP_TLS_CLIENT_CA_CERT_PATH",
1440 Some(cert.path().to_str().unwrap()),
1441 ),
1442 ],
1443 || assert!(TcpStreamServer::build_tls_acceptor().is_err()),
1444 );
1445 }
1446
1447 struct FailingIpResolver;
1449
1450 impl IpResolver for FailingIpResolver {
1451 fn local_ip(&self) -> Result<std::net::IpAddr, Error> {
1452 Err(Error::LocalIpAddressNotFound)
1453 }
1454
1455 fn local_ipv6(&self) -> Result<std::net::IpAddr, Error> {
1456 Err(Error::LocalIpAddressNotFound)
1457 }
1458 }
1459
1460 #[tokio::test]
1461 async fn test_tcp_stream_server_default_behavior() {
1462 let options = ServerOptions::default();
1465 let result = TcpStreamServer::new(options).await;
1466
1467 assert!(
1468 result.is_ok(),
1469 "TcpStreamServer::new should succeed with default options"
1470 );
1471
1472 let server = result.unwrap();
1473
1474 let context = Context::new(());
1476 let stream_options = StreamOptions::builder()
1477 .context(context.context())
1478 .enable_request_stream(false)
1479 .enable_response_stream(true)
1480 .build()
1481 .unwrap();
1482
1483 let pending_connection = server.register(stream_options).await;
1484
1485 let connection_info = pending_connection
1487 .recv_stream
1488 .as_ref()
1489 .unwrap()
1490 .connection_info
1491 .clone();
1492
1493 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1494 let socket_addr = tcp_info.address.parse::<std::net::SocketAddr>().unwrap();
1495
1496 assert!(
1498 socket_addr.port() > 0,
1499 "Server should be assigned a valid port number"
1500 );
1501
1502 println!(
1503 "Server created successfully with address: {}",
1504 tcp_info.address
1505 );
1506 }
1507
1508 #[test]
1514 fn data_plane_channel_capacity_matches_send_buffer_count() {
1515 let (tx, _rx) = data_plane_channel::<()>(7);
1516 assert_eq!(tx.max_capacity(), 7);
1517
1518 let (tx, _rx) = data_plane_channel::<()>(DEFAULT_SEND_BUFFER_COUNT);
1519 assert_eq!(tx.max_capacity(), 64);
1520
1521 let (tx, _rx) = data_plane_channel::<()>(0);
1523 assert_eq!(tx.max_capacity(), 1);
1524 }
1525
1526 #[tokio::test]
1531 async fn register_threads_send_buffer_count_into_connection_structs() {
1532 let server = TcpStreamServer::new(ServerOptions::default())
1533 .await
1534 .expect("server");
1535 let context = Context::new(());
1536 let options = StreamOptions::builder()
1537 .context(context.context())
1538 .enable_request_stream(true)
1539 .enable_response_stream(true)
1540 .send_buffer_count(7)
1541 .build()
1542 .unwrap();
1543
1544 let _pending = server.register(options).await;
1545
1546 let state = server.state.lock();
1547 assert_eq!(state.tx_subjects.len(), 1, "one request stream registered");
1548 assert_eq!(state.rx_subjects.len(), 1, "one response stream registered");
1549 assert!(
1550 state.tx_subjects.values().all(|c| c.send_buffer_count == 7),
1551 "send_buffer_count must reach RequestedSendConnection"
1552 );
1553 assert!(
1554 state.rx_subjects.values().all(|c| c.send_buffer_count == 7),
1555 "send_buffer_count must reach RequestedRecvConnection"
1556 );
1557 }
1558
1559 #[tokio::test]
1560 async fn test_tcp_stream_server_fallback_to_loopback() {
1561 let options = ServerOptions::builder().port(0).build().unwrap();
1565
1566 let result = TcpStreamServer::new_with_resolver(options, FailingIpResolver).await;
1568 assert!(
1569 result.is_ok(),
1570 "Server creation should succeed with fallback even when IP detection fails"
1571 );
1572
1573 let server = result.unwrap();
1574
1575 let context = Context::new(());
1577 let stream_options = StreamOptions::builder()
1578 .context(context.context())
1579 .enable_request_stream(false)
1580 .enable_response_stream(true)
1581 .build()
1582 .unwrap();
1583
1584 let pending_connection = server.register(stream_options).await;
1585 let connection_info = pending_connection
1586 .recv_stream
1587 .as_ref()
1588 .unwrap()
1589 .connection_info
1590 .clone();
1591
1592 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1593 let socket_addr = tcp_info.address.parse::<std::net::SocketAddr>().unwrap();
1594
1595 let ip = socket_addr.ip();
1597 assert!(
1598 ip.is_loopback(),
1599 "Should use loopback when IP detection fails"
1600 );
1601
1602 assert_eq!(
1604 ip,
1605 std::net::IpAddr::V4(std::net::Ipv4Addr::new(127, 0, 0, 1)),
1606 "Fallback should use exactly 127.0.0.1, got: {}",
1607 ip
1608 );
1609
1610 println!("SUCCESS: Fallback to 127.0.0.1 was confirmed: {}", ip);
1611
1612 assert!(socket_addr.port() > 0, "Server should have a valid port");
1614 }
1615
1616 async fn test_server() -> Arc<TcpStreamServer> {
1618 TcpStreamServer::new_with_resolver(
1619 ServerOptions::builder().port(0).build().unwrap(),
1620 FailingIpResolver,
1621 )
1622 .await
1623 .unwrap()
1624 }
1625
1626 async fn register_and_get_subject(
1628 server: &TcpStreamServer,
1629 ) -> (
1630 String,
1631 tokio::sync::oneshot::Receiver<Result<super::StreamReceiver, StreamPrologueError>>,
1632 ) {
1633 let context = Context::new(());
1634 let options = StreamOptions::builder()
1635 .context(context.context())
1636 .enable_request_stream(false)
1637 .enable_response_stream(true)
1638 .build()
1639 .unwrap();
1640
1641 let pending = server.register(options).await;
1642 let recv_stream = pending.recv_stream.unwrap();
1643 let (conn_info, provider) = recv_stream.into_parts();
1644 let tcp_info: TcpStreamConnectionInfo = conn_info.try_into().unwrap();
1645 (tcp_info.subject, provider)
1646 }
1647
1648 fn make_eid(
1650 namespace: &str,
1651 component: &str,
1652 endpoint: &str,
1653 instance_id: u64,
1654 ) -> EndpointInstanceId {
1655 EndpointInstanceId {
1656 namespace: namespace.to_string(),
1657 component: component.to_string(),
1658 endpoint: endpoint.to_string(),
1659 instance_id,
1660 }
1661 }
1662
1663 async fn register_and_get_bidi_subjects(
1666 server: &TcpStreamServer,
1667 ) -> (
1668 String,
1669 tokio::sync::oneshot::Receiver<Result<super::StreamSender, StreamPrologueError>>,
1670 String,
1671 tokio::sync::oneshot::Receiver<Result<super::StreamReceiver, StreamPrologueError>>,
1672 ) {
1673 let context = Context::new(());
1674 let options = StreamOptions::builder()
1675 .context(context.context())
1676 .enable_request_stream(true)
1677 .enable_response_stream(true)
1678 .build()
1679 .unwrap();
1680
1681 let pending = server.register(options).await;
1682 let send_stream = pending.send_stream.unwrap();
1683 let recv_stream = pending.recv_stream.unwrap();
1684 let (send_info, send_provider) = send_stream.into_parts();
1685 let (recv_info, recv_provider) = recv_stream.into_parts();
1686 let send_tcp_info: TcpStreamConnectionInfo = send_info.try_into().unwrap();
1687 let recv_tcp_info: TcpStreamConnectionInfo = recv_info.try_into().unwrap();
1688 (
1689 send_tcp_info.subject,
1690 send_provider,
1691 recv_tcp_info.subject,
1692 recv_provider,
1693 )
1694 }
1695
1696 #[tokio::test]
1701 async fn test_cancel_instance_streams_drops_both_bidi_halves() {
1702 let server = test_server().await;
1703 let (send_subj, send_provider, recv_subj, recv_provider) =
1704 register_and_get_bidi_subjects(&server).await;
1705
1706 let id = make_eid("ns", "comp", "generate", 7);
1707 assert!(
1708 server
1709 .associate_instance(&recv_subj, Some(&send_subj), &id)
1710 .await,
1711 "fresh instance must not be tombstoned"
1712 );
1713
1714 let cancelled = server.cancel_instance_streams(&id).await;
1715 assert_eq!(cancelled, 2, "both request + response halves must count");
1716
1717 assert!(
1718 recv_provider.await.is_err(),
1719 "recv provider should resolve with RecvError"
1720 );
1721 assert!(
1722 send_provider.await.is_err(),
1723 "send provider should resolve with RecvError after instance cancellation"
1724 );
1725 }
1726
1727 #[tokio::test]
1730 async fn test_associate_instance_tombstone_cancels_both_bidi_halves() {
1731 let server = test_server().await;
1732 let id = make_eid("ns", "comp", "generate", 8);
1733 server.cancel_instance_streams(&id).await;
1735
1736 let (send_subj, send_provider, recv_subj, recv_provider) =
1737 register_and_get_bidi_subjects(&server).await;
1738
1739 assert!(
1740 !server
1741 .associate_instance(&recv_subj, Some(&send_subj), &id)
1742 .await,
1743 "tombstoned instance must reject association"
1744 );
1745
1746 assert!(recv_provider.await.is_err());
1747 assert!(send_provider.await.is_err());
1748 }
1749
1750 #[tokio::test]
1751 async fn test_cancel_instance_streams_unblocks_receiver() {
1752 let server = test_server().await;
1753
1754 let (subject, provider) = register_and_get_subject(&server).await;
1755
1756 let id = make_eid("ns", "comp", "generate", 42);
1757 assert!(server.associate_instance(&subject, None, &id).await);
1758
1759 let cancelled = server.cancel_instance_streams(&id).await;
1760 assert_eq!(cancelled, 1);
1761
1762 let result = provider.await;
1764 assert!(result.is_err(), "Expected RecvError after cancellation");
1765 }
1766
1767 #[tokio::test]
1768 async fn test_cancel_instance_streams_multiple_subjects() {
1769 let server = test_server().await;
1770
1771 let (subj1, prov1) = register_and_get_subject(&server).await;
1772 let (subj2, prov2) = register_and_get_subject(&server).await;
1773 let (subj3, prov3) = register_and_get_subject(&server).await;
1774
1775 let id10 = make_eid("ns", "comp", "generate", 10);
1776 let id20 = make_eid("ns", "comp", "generate", 20);
1777
1778 assert!(server.associate_instance(&subj1, None, &id10).await);
1780 assert!(server.associate_instance(&subj2, None, &id10).await);
1781 assert!(server.associate_instance(&subj3, None, &id20).await);
1782
1783 let cancelled = server.cancel_instance_streams(&id10).await;
1785 assert_eq!(cancelled, 2);
1786
1787 assert!(prov1.await.is_err());
1788 assert!(prov2.await.is_err());
1789
1790 let cancelled = server.cancel_instance_streams(&id20).await;
1792 assert_eq!(cancelled, 1);
1793 assert!(prov3.await.is_err());
1794 }
1795
1796 #[tokio::test]
1797 async fn test_cancel_instance_streams_nonexistent_instance() {
1798 let server = test_server().await;
1799
1800 let id = make_eid("ns", "comp", "generate", 999);
1801 let cancelled = server.cancel_instance_streams(&id).await;
1802 assert_eq!(cancelled, 0);
1803 }
1804
1805 #[tokio::test]
1806 async fn test_cancel_recv_stream_cleans_up_instance_tracking() {
1807 let server = test_server().await;
1808
1809 let (subject, _provider) = register_and_get_subject(&server).await;
1810 let id = make_eid("ns", "comp", "generate", 42);
1811 assert!(server.associate_instance(&subject, None, &id).await);
1812
1813 server.cancel_recv_stream(&subject).await;
1815
1816 let cancelled = server.cancel_instance_streams(&id).await;
1818 assert_eq!(
1819 cancelled, 0,
1820 "Instance tracking should have been cleaned up"
1821 );
1822 }
1823
1824 #[tokio::test]
1825 async fn test_registered_stream_drop_runs_cleanup() {
1826 let server = test_server().await;
1827
1828 let context = Context::new(());
1830 let options = StreamOptions::builder()
1831 .context(context.context())
1832 .enable_request_stream(false)
1833 .enable_response_stream(true)
1834 .build()
1835 .unwrap();
1836
1837 let pending = server.register(options).await;
1838 let recv_stream = pending.recv_stream.unwrap();
1839
1840 let tcp_info: TcpStreamConnectionInfo =
1842 recv_stream.connection_info.clone().try_into().unwrap();
1843 let subject = tcp_info.subject.clone();
1844
1845 {
1847 let state = server.state.lock();
1848 assert!(state.rx_subjects.contains_key(&subject));
1849 }
1850
1851 drop(recv_stream);
1853
1854 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1856
1857 {
1859 let state = server.state.lock();
1860 assert!(
1861 !state.rx_subjects.contains_key(&subject),
1862 "RAII cleanup should have removed the rx_subjects entry"
1863 );
1864 }
1865 }
1866
1867 #[tokio::test]
1868 async fn test_registered_stream_into_parts_disarms_cleanup() {
1869 let server = test_server().await;
1870
1871 let context = Context::new(());
1872 let options = StreamOptions::builder()
1873 .context(context.context())
1874 .enable_request_stream(false)
1875 .enable_response_stream(true)
1876 .build()
1877 .unwrap();
1878
1879 let pending = server.register(options).await;
1880 let recv_stream = pending.recv_stream.unwrap();
1881
1882 let tcp_info: TcpStreamConnectionInfo =
1883 recv_stream.connection_info.clone().try_into().unwrap();
1884 let subject = tcp_info.subject.clone();
1885
1886 let (_conn_info, _provider) = recv_stream.into_parts();
1888
1889 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1891
1892 {
1894 let state = server.state.lock();
1895 assert!(
1896 state.rx_subjects.contains_key(&subject),
1897 "into_parts() should disarm the RAII cleanup"
1898 );
1899 }
1900 }
1901
1902 #[tokio::test]
1903 async fn test_associate_after_cancel_is_immediately_cancelled() {
1904 let server = test_server().await;
1906
1907 let id = make_eid("ns", "comp", "generate", 42);
1908
1909 let cancelled = server.cancel_instance_streams(&id).await;
1911 assert_eq!(cancelled, 0);
1912
1913 let (subject, provider) = register_and_get_subject(&server).await;
1915 let associated = server.associate_instance(&subject, None, &id).await;
1916
1917 assert!(
1919 !associated,
1920 "associate_instance on a tombstoned instance should return false"
1921 );
1922
1923 let result = provider.await;
1926 assert!(
1927 result.is_err(),
1928 "Late associate_instance on a tombstoned instance should immediately cancel"
1929 );
1930 }
1931
1932 #[tokio::test]
1933 async fn test_clear_tombstone_allows_new_associations() {
1934 let server = test_server().await;
1935
1936 let id = make_eid("ns", "comp", "generate", 42);
1937
1938 server.cancel_instance_streams(&id).await;
1939 server.clear_instance_tombstone(&id).await;
1940
1941 let (subject, _provider) = register_and_get_subject(&server).await;
1943 assert!(server.associate_instance(&subject, None, &id).await);
1944
1945 let cancelled = server.cancel_instance_streams(&id).await;
1947 assert_eq!(
1948 cancelled, 1,
1949 "After clearing tombstone, subjects should be tracked normally"
1950 );
1951 }
1952
1953 #[tokio::test]
1954 async fn test_cancel_does_not_affect_sibling_endpoint() {
1955 let server = test_server().await;
1958
1959 let (gen_subj, gen_prov) = register_and_get_subject(&server).await;
1960 let (pre_subj, pre_prov) = register_and_get_subject(&server).await;
1961
1962 let gen_id = make_eid("ns", "comp", "generate", 42);
1963 let pre_id = make_eid("ns", "comp", "prefill", 42);
1964
1965 assert!(server.associate_instance(&gen_subj, None, &gen_id).await);
1966 assert!(server.associate_instance(&pre_subj, None, &pre_id).await);
1967
1968 let cancelled = server.cancel_instance_streams(&gen_id).await;
1970 assert_eq!(
1971 cancelled, 1,
1972 "Only the generate subject should be cancelled"
1973 );
1974 assert!(gen_prov.await.is_err());
1975
1976 let still_pending = server.cancel_instance_streams(&pre_id).await;
1978 assert_eq!(still_pending, 1, "prefill subject should still be tracked");
1979 assert!(pre_prov.await.is_err());
1980 }
1981
1982 #[tokio::test]
1983 async fn test_tombstone_is_endpoint_scoped() {
1984 let server = test_server().await;
1987
1988 let gen_id = make_eid("ns", "comp", "generate", 42);
1989 let pre_id = make_eid("ns", "comp", "prefill", 42);
1990
1991 server.cancel_instance_streams(&gen_id).await;
1992
1993 let (gen_subj, gen_prov) = register_and_get_subject(&server).await;
1995 assert!(
1996 !server.associate_instance(&gen_subj, None, &gen_id).await,
1997 "generate should be tombstoned"
1998 );
1999 assert!(gen_prov.await.is_err());
2000
2001 let (pre_subj, _pre_prov) = register_and_get_subject(&server).await;
2003 assert!(
2004 server.associate_instance(&pre_subj, None, &pre_id).await,
2005 "prefill tombstone is independent; subject should be tracked"
2006 );
2007 let count = server.cancel_instance_streams(&pre_id).await;
2008 assert_eq!(count, 1, "prefill subject should be tracked normally");
2009 }
2010
2011 #[tokio::test]
2012 async fn test_cancel_does_not_affect_different_component() {
2013 let server = test_server().await;
2017
2018 let (subj_a, prov_a) = register_and_get_subject(&server).await;
2019 let (subj_b, prov_b) = register_and_get_subject(&server).await;
2020
2021 let id_a = make_eid("ns-a", "comp-a", "generate", 42);
2023 let id_b = make_eid("ns-b", "comp-b", "generate", 42);
2024
2025 assert!(server.associate_instance(&subj_a, None, &id_a).await);
2026 assert!(server.associate_instance(&subj_b, None, &id_b).await);
2027
2028 let cancelled = server.cancel_instance_streams(&id_a).await;
2030 assert_eq!(cancelled, 1, "Only service-A subject should be cancelled");
2031 assert!(prov_a.await.is_err());
2032
2033 let still_tracked = server.cancel_instance_streams(&id_b).await;
2035 assert_eq!(still_tracked, 1, "Service-B subject should be unaffected");
2036 assert!(prov_b.await.is_err());
2037 }
2038
2039 #[tokio::test(start_paused = true)]
2040 async fn test_tombstone_expires_after_ttl() {
2041 let server = test_server().await;
2045
2046 let id = make_eid("ns", "comp", "generate", 42);
2047
2048 server.cancel_instance_streams(&id).await;
2050 {
2051 let state = server.state.lock();
2052 assert!(state.removed_instances.contains_key(&id));
2053 }
2054
2055 tokio::time::advance(TOMBSTONE_TTL + Duration::from_secs(1)).await;
2057
2058 let (subject, _provider) = register_and_get_subject(&server).await;
2061 assert!(
2062 server.associate_instance(&subject, None, &id).await,
2063 "tombstone older than TTL should not block association"
2064 );
2065
2066 {
2069 let state = server.state.lock();
2070 assert!(
2071 !state.removed_instances.contains_key(&id),
2072 "expired tombstone should be pruned, not retained"
2073 );
2074 }
2075 }
2076
2077 #[tokio::test(start_paused = true)]
2078 async fn test_tombstone_within_ttl_blocks_associate() {
2079 let server = test_server().await;
2082
2083 let id = make_eid("ns", "comp", "generate", 42);
2084 server.cancel_instance_streams(&id).await;
2085
2086 tokio::time::advance(Duration::from_secs(1)).await;
2088
2089 let (subject, provider) = register_and_get_subject(&server).await;
2090 assert!(
2091 !server.associate_instance(&subject, None, &id).await,
2092 "tombstone within TTL must still block association"
2093 );
2094 assert!(provider.await.is_err());
2095 }
2096
2097 #[tokio::test(start_paused = true)]
2098 async fn test_tombstone_lazy_prune_on_cancel() {
2099 let server = test_server().await;
2102
2103 let id_old = make_eid("ns", "comp", "generate", 1);
2104 let id_new = make_eid("ns", "comp", "generate", 2);
2105
2106 server.cancel_instance_streams(&id_old).await;
2107 tokio::time::advance(TOMBSTONE_TTL + Duration::from_secs(1)).await;
2108 server.cancel_instance_streams(&id_new).await;
2109
2110 let state = server.state.lock();
2111 assert!(
2112 !state.removed_instances.contains_key(&id_old),
2113 "old tombstone should be pruned by the next cancel_instance_streams call"
2114 );
2115 assert!(
2116 state.removed_instances.contains_key(&id_new),
2117 "fresh tombstone should be retained"
2118 );
2119 assert_eq!(state.removed_instances.len(), 1);
2120 }
2121
2122 #[tokio::test]
2123 async fn test_clear_tombstone_only_affects_named_identity() {
2124 let server = test_server().await;
2129
2130 let id_a = make_eid("ns", "comp", "generate", 1);
2131 let id_b = make_eid("ns", "comp", "generate", 2);
2132
2133 server.cancel_instance_streams(&id_a).await;
2134 server.clear_instance_tombstone(&id_b).await;
2135
2136 let state = server.state.lock();
2137 assert!(
2138 state.removed_instances.contains_key(&id_a),
2139 "clearing a different identity must not remove id_a's tombstone"
2140 );
2141 }
2142
2143 #[tokio::test]
2144 async fn test_tombstone_scoped_to_full_identity() {
2145 let server = test_server().await;
2148
2149 let id_a = make_eid("ns-a", "comp-a", "generate", 42);
2150 let id_b = make_eid("ns-b", "comp-b", "generate", 42);
2151
2152 server.cancel_instance_streams(&id_a).await;
2154
2155 let (subj_a, prov_a) = register_and_get_subject(&server).await;
2157 assert!(!server.associate_instance(&subj_a, None, &id_a).await);
2158 assert!(prov_a.await.is_err());
2159
2160 let (subj_b, _prov_b) = register_and_get_subject(&server).await;
2162 assert!(
2163 server.associate_instance(&subj_b, None, &id_b).await,
2164 "Different namespace/component must not be tombstoned"
2165 );
2166 assert_eq!(server.cancel_instance_streams(&id_b).await, 1);
2167 }
2168
2169 type TestFramedRead = FramedRead<ReadHalf<TcpStream>, TwoPartCodec>;
2170 type TestFramedWrite = FramedWrite<WriteHalf<TcpStream>, TwoPartCodec>;
2171 type TestResponseStream = (TestFramedRead, TestFramedWrite, StreamReceiver);
2172
2173 async fn open_registered_response_stream() -> TestResponseStream {
2177 let options = ServerOptions::builder().port(0).build().unwrap();
2178 let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver)
2179 .await
2180 .unwrap();
2181 let context = Context::new(());
2182 let stream_options = StreamOptions::builder()
2183 .context(context.context())
2184 .enable_request_stream(false)
2185 .enable_response_stream(true)
2186 .build()
2187 .unwrap();
2188 let pending_connection = server.register(stream_options).await;
2189 let registered_stream = pending_connection.recv_stream.unwrap();
2190 let (connection_info, stream_provider) = registered_stream.into_parts();
2191 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
2192
2193 let stream = TcpStream::connect(&tcp_info.address).await.unwrap();
2194 let (read_half, write_half) = tokio::io::split(stream);
2195 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
2196 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
2197
2198 let handshake = CallHomeHandshake {
2199 subject: tcp_info.subject,
2200 stream_type: StreamType::Response,
2201 };
2202 framed_writer
2203 .send(TwoPartMessage::from_header(
2204 serde_json::to_vec(&handshake).unwrap().into(),
2205 ))
2206 .await
2207 .unwrap();
2208 framed_writer
2209 .send(TwoPartMessage::from_header(
2210 serde_json::to_vec(&ResponseStreamPrologue {
2211 error: None,
2212 typed_error: None,
2213 })
2214 .unwrap()
2215 .into(),
2216 ))
2217 .await
2218 .unwrap();
2219
2220 let receiver = tokio::time::timeout(std::time::Duration::from_secs(1), stream_provider)
2223 .await
2224 .expect("server should establish response stream within timeout")
2225 .expect("stream provider should not be dropped")
2226 .expect("response stream should be accepted");
2227
2228 (framed_reader, framed_writer, receiver)
2229 }
2230
2231 async fn recv_control_message(framed_reader: &mut TestFramedRead) -> ControlMessage {
2232 let message = tokio::time::timeout(std::time::Duration::from_secs(1), framed_reader.next())
2235 .await
2236 .expect("server should send a control message within timeout")
2237 .expect("server should not close before sending control")
2238 .expect("control message should decode");
2239 let (header, data) = message.optional_parts();
2240 assert!(data.is_none(), "control message should not contain data");
2241 serde_json::from_slice(header.expect("control header missing").as_ref()).unwrap()
2242 }
2243
2244 #[tokio::test]
2249 async fn test_tcp_stream_server_sends_kill_on_unexpected_control_message() {
2250 let (mut framed_reader, mut framed_writer, _receiver) =
2251 open_registered_response_stream().await;
2252
2253 framed_writer
2254 .send(TwoPartMessage::from_header(
2255 serde_json::to_vec(&ControlMessage::Stop).unwrap().into(),
2256 ))
2257 .await
2258 .unwrap();
2259
2260 assert_eq!(
2261 recv_control_message(&mut framed_reader).await,
2262 ControlMessage::Kill,
2263 "unexpected control message should kill only this stream"
2264 );
2265 }
2266
2267 #[tokio::test]
2271 async fn test_tcp_stream_server_sends_kill_on_read_error() {
2272 let (mut framed_reader, framed_writer, _receiver) = open_registered_response_stream().await;
2273
2274 let mut raw_writer = framed_writer.into_inner();
2275 raw_writer.write_all(&[0u8; 8]).await.unwrap();
2276 raw_writer.shutdown().await.unwrap();
2277
2278 assert_eq!(
2279 recv_control_message(&mut framed_reader).await,
2280 ControlMessage::Kill,
2281 "framing read error should kill only this stream"
2282 );
2283 }
2284
2285 #[tokio::test]
2288 async fn test_tcp_stream_server_sends_kill_on_sentinel_with_data() {
2289 let (mut framed_reader, mut framed_writer, _receiver) =
2290 open_registered_response_stream().await;
2291
2292 let header = serde_json::to_vec(&ControlMessage::Sentinel)
2293 .unwrap()
2294 .into();
2295 framed_writer
2296 .send(TwoPartMessage::from_parts(
2297 header,
2298 Bytes::from_static(b"unexpected payload"),
2299 ))
2300 .await
2301 .unwrap();
2302
2303 assert_eq!(
2304 recv_control_message(&mut framed_reader).await,
2305 ControlMessage::Kill,
2306 "Sentinel with data should kill only this stream"
2307 );
2308 }
2309
2310 #[tokio::test]
2314 async fn test_tcp_stream_server_returns_error_on_invalid_prologue() {
2315 let options = ServerOptions::builder().port(0).build().unwrap();
2316 let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver)
2317 .await
2318 .unwrap();
2319 let context = Context::new(());
2320 let stream_options = StreamOptions::builder()
2321 .context(context.context())
2322 .enable_request_stream(false)
2323 .enable_response_stream(true)
2324 .build()
2325 .unwrap();
2326 let pending_connection = server.register(stream_options).await;
2327 let registered_stream = pending_connection.recv_stream.unwrap();
2328 let (connection_info, stream_provider) = registered_stream.into_parts();
2329 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
2330
2331 let stream = TcpStream::connect(&tcp_info.address).await.unwrap();
2332 let (_read_half, write_half) = tokio::io::split(stream);
2333 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
2334
2335 let handshake = CallHomeHandshake {
2336 subject: tcp_info.subject,
2337 stream_type: StreamType::Response,
2338 };
2339 framed_writer
2340 .send(TwoPartMessage::from_header(
2341 serde_json::to_vec(&handshake).unwrap().into(),
2342 ))
2343 .await
2344 .unwrap();
2345
2346 framed_writer
2348 .send(TwoPartMessage::from_data(Bytes::from_static(
2349 b"not a prologue",
2350 )))
2351 .await
2352 .unwrap();
2353
2354 let outcome = tokio::time::timeout(std::time::Duration::from_secs(1), stream_provider)
2355 .await
2356 .expect("stream provider should resolve quickly")
2357 .expect("stream provider channel should not be dropped");
2358 match outcome {
2360 Err(err) => assert!(
2361 err.contains("malformed prologue"),
2362 "expected malformed-prologue error, got: {err}"
2363 ),
2364 Ok(_) => panic!("invalid prologue should produce an error, but got Ok"),
2365 }
2366 }
2367
2368 #[tokio::test]
2371 async fn test_unknown_typed_error_preserves_the_legacy_prologue_error() {
2372 let options = ServerOptions::builder().port(0).build().unwrap();
2373 let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver)
2374 .await
2375 .unwrap();
2376 let context = Context::new(());
2377 let stream_options = StreamOptions::builder()
2378 .context(context.context())
2379 .enable_request_stream(false)
2380 .enable_response_stream(true)
2381 .build()
2382 .unwrap();
2383 let pending_connection = server.register(stream_options).await;
2384 let (connection_info, stream_provider) =
2385 pending_connection.recv_stream.unwrap().into_parts();
2386 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
2387
2388 let stream = TcpStream::connect(&tcp_info.address).await.unwrap();
2389 let (_read_half, write_half) = tokio::io::split(stream);
2390 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
2391
2392 let handshake = CallHomeHandshake {
2393 subject: tcp_info.subject,
2394 stream_type: StreamType::Response,
2395 };
2396 framed_writer
2397 .send(TwoPartMessage::from_header(
2398 serde_json::to_vec(&handshake).unwrap().into(),
2399 ))
2400 .await
2401 .unwrap();
2402
2403 framed_writer
2406 .send(TwoPartMessage::from_header(Bytes::from_static(
2407 br#"{"error":"Generate Error: boom","typed_error":{"error_type":"VariantFromTheFuture","message":"boom"}}"#,
2408 )))
2409 .await
2410 .unwrap();
2411
2412 let outcome = tokio::time::timeout(std::time::Duration::from_secs(5), stream_provider)
2413 .await
2414 .expect("stream provider should resolve quickly")
2415 .expect("the oneshot must be notified, not dropped");
2416
2417 match outcome {
2419 Err(err) => {
2420 assert_eq!(err.message, "Generate Error: boom");
2421 assert!(err.typed_error.is_none());
2422 }
2423 Ok(_) => panic!("an error prologue must not yield a usable stream"),
2424 }
2425 }
2426
2427 #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
2428 async fn test_concurrent_response_registration_and_call_home() {
2429 const STREAMS: usize = 128;
2430
2431 let result = time::timeout(Duration::from_secs(20), async {
2432 let server = test_server().await;
2433 let mut pending_streams = Vec::with_capacity(STREAMS);
2434 let mut client_tasks = Vec::with_capacity(STREAMS);
2435
2436 for idx in 0..STREAMS {
2437 let context = Context::new(());
2438 let options = StreamOptions::builder()
2439 .context(context.context())
2440 .enable_request_stream(false)
2441 .enable_response_stream(true)
2442 .build()
2443 .unwrap();
2444
2445 let pending = server.register(options).await;
2446 let registered_stream = pending.recv_stream.unwrap();
2447 let (connection_info, stream_provider) = registered_stream.into_parts();
2448 let client_context =
2449 Context::with_id_and_metadata((), context.id().to_string(), Default::default());
2450 let payload = Bytes::from(format!("payload-{idx}"));
2451
2452 pending_streams.push((idx, payload.clone(), stream_provider));
2453 client_tasks.push(tokio::spawn(async move {
2454 let mut sender = TcpClient::create_response_stream(
2455 client_context.context(),
2456 connection_info,
2457 None,
2458 )
2459 .await
2460 .unwrap();
2461 sender.send_prologue(None).await.unwrap();
2462 sender.send(payload).await.unwrap();
2463 }));
2464 }
2465
2466 for task in client_tasks {
2467 task.await.unwrap();
2468 }
2469
2470 for (idx, expected, stream_provider) in pending_streams {
2471 let mut stream = stream_provider.await.unwrap().unwrap();
2472 let actual = stream.rx.recv().await.unwrap();
2473 assert_eq!(actual, expected, "payload mismatch for stream {idx}");
2474 }
2475 })
2476 .await;
2477
2478 assert!(
2479 result.is_ok(),
2480 "concurrent response registration and call-home timed out"
2481 );
2482 }
2483
2484 #[tokio::test]
2487 async fn test_typed_prologue_error_survives_to_requester() {
2488 let server = test_server().await;
2489 let context = Context::new(());
2490 let options = StreamOptions::builder()
2491 .context(context.context())
2492 .enable_request_stream(false)
2493 .enable_response_stream(true)
2494 .build()
2495 .unwrap();
2496
2497 let pending = server.register(options).await;
2498 let (connection_info, stream_provider) = pending.recv_stream.unwrap().into_parts();
2499 let client_context =
2500 Context::with_id_and_metadata((), context.id().to_string(), Default::default());
2501
2502 let worker_error = DynamoError::builder()
2503 .error_type(ErrorType::Backend(BackendError::InvalidArgument))
2504 .message("multimodal input is not supported by this backend")
2505 .build();
2506
2507 let mut sender =
2508 TcpClient::create_response_stream(client_context.context(), connection_info, None)
2509 .await
2510 .unwrap();
2511 sender
2512 .send_prologue_typed(Some(StreamPrologueError::new(
2513 "Generate Error: multimodal input is not supported by this backend",
2514 worker_error,
2515 )))
2516 .await
2517 .unwrap();
2518
2519 let outcome = tokio::time::timeout(std::time::Duration::from_secs(5), stream_provider)
2520 .await
2521 .expect("stream provider should resolve quickly")
2522 .expect("stream provider channel should not be dropped");
2523
2524 let prologue_error = match outcome {
2526 Err(err) => err,
2527 Ok(_) => panic!("an error prologue must not yield a usable stream"),
2528 };
2529 assert_eq!(
2530 prologue_error.typed_error.as_ref().map(|e| e.error_type()),
2531 Some(ErrorType::Backend(BackendError::InvalidArgument)),
2532 "the worker's error type must survive the prologue round trip"
2533 );
2534 }
2535
2536 use futures::SinkExt;
2544
2545 async fn register_and_dial_request_stream(
2550 server: &TcpStreamServer,
2551 ) -> (
2552 FramedRead<tokio::io::ReadHalf<TcpStream>, TwoPartCodec>,
2553 super::StreamSender,
2554 Arc<dyn AsyncEngineContext>,
2555 ) {
2556 let upstream_ctx = Context::new(()).context();
2557 let options = StreamOptions::builder()
2558 .context(upstream_ctx.clone())
2559 .enable_request_stream(true)
2560 .enable_response_stream(false)
2561 .build()
2562 .unwrap();
2563
2564 let pending = server.register(options).await;
2565 let send_stream = pending.send_stream.unwrap();
2566 let (conn_info, send_provider) = send_stream.into_parts();
2567 let tcp_info: TcpStreamConnectionInfo = conn_info.try_into().unwrap();
2568
2569 let raw = TcpStream::connect(&tcp_info.address).await.unwrap();
2570 let (read_half, write_half) = tokio::io::split(raw);
2571 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
2572 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
2573
2574 let handshake = super::CallHomeHandshake {
2575 subject: tcp_info.subject.clone(),
2576 stream_type: StreamType::Request,
2577 };
2578 let handshake_bytes = serde_json::to_vec(&handshake).unwrap();
2579 framed_writer
2580 .send(TwoPartMessage::from_header(handshake_bytes.into()))
2581 .await
2582 .unwrap();
2583 drop(framed_writer);
2584
2585 let sender = send_provider.await.unwrap().unwrap();
2586 (framed_reader, sender, upstream_ctx)
2587 }
2588
2589 async fn next_control_message(
2592 reader: &mut FramedRead<tokio::io::ReadHalf<TcpStream>, TwoPartCodec>,
2593 ) -> ControlMessage {
2594 loop {
2595 let frame = reader
2596 .next()
2597 .await
2598 .expect("socket closed before control message arrived")
2599 .expect("decode error");
2600 if let Some(header) = frame.header() {
2601 return serde_json::from_slice::<ControlMessage>(header)
2602 .expect("invalid control message bytes");
2603 }
2604 }
2606 }
2607
2608 #[tokio::test]
2611 async fn test_request_stream_sends_sentinel_on_clean_drop() {
2612 let server = test_server().await;
2613 let (mut reader, sender, _ctx) = register_and_dial_request_stream(&server).await;
2614
2615 drop(sender);
2616
2617 let ctrl = next_control_message(&mut reader).await;
2618 assert!(
2619 matches!(ctrl, ControlMessage::Sentinel),
2620 "clean drain should emit Sentinel, got {ctrl:?}"
2621 );
2622 }
2623
2624 #[tokio::test]
2627 async fn test_request_stream_sends_kill_on_context_killed() {
2628 let server = test_server().await;
2629 let (mut reader, _sender, ctx) = register_and_dial_request_stream(&server).await;
2630
2631 ctx.kill();
2632
2633 let ctrl = next_control_message(&mut reader).await;
2634 assert!(
2635 matches!(ctrl, ControlMessage::Kill),
2636 "context.kill() should emit Kill, got {ctrl:?}"
2637 );
2638 }
2639
2640 #[tokio::test]
2643 async fn test_request_stream_sends_stop_on_context_stopped() {
2644 let server = test_server().await;
2645 let (mut reader, _sender, ctx) = register_and_dial_request_stream(&server).await;
2646
2647 ctx.stop();
2648
2649 let ctrl = next_control_message(&mut reader).await;
2650 assert!(
2651 matches!(ctrl, ControlMessage::Stop),
2652 "context.stop() should emit Stop, got {ctrl:?}"
2653 );
2654 }
2655
2656 #[cfg(unix)]
2659 fn emfile_error() -> std::io::Error {
2660 std::io::Error::from_raw_os_error(libc::EMFILE)
2661 }
2662
2663 #[cfg(unix)]
2664 #[test]
2665 fn accept_backoff_classifies_all_exhaustion_errnos() {
2666 for errno in [libc::EMFILE, libc::ENFILE, libc::ENOBUFS, libc::ENOMEM] {
2667 let err = std::io::Error::from_raw_os_error(errno);
2668 assert_eq!(
2669 AcceptBackoff::classify(&err),
2670 AcceptFailure::Exhaustion,
2671 "errno {errno} must classify as exhaustion"
2672 );
2673 }
2674 let ordinary = std::io::Error::from(std::io::ErrorKind::ConnectionAborted);
2676 assert_eq!(AcceptBackoff::classify(&ordinary), AcceptFailure::Ordinary,);
2677 }
2678
2679 #[cfg(unix)]
2680 #[test]
2681 fn accept_backoff_grows_resets_and_rate_limits() {
2682 let mut backoff = AcceptBackoff::default();
2683 let now = std::time::Instant::now();
2684
2685 let delays: Vec<Duration> = (0..12)
2687 .map(|_| backoff.record_exhaustion(now).delay)
2688 .collect();
2689 assert!(delays[0] > Duration::ZERO);
2690 let sat = delays
2691 .iter()
2692 .position(|d| *d == ACCEPT_BACKOFF_MAX_DELAY)
2693 .expect("delay reaches the ceiling");
2694 for pair in delays[..=sat].windows(2) {
2695 assert!(
2696 pair[1] > pair[0],
2697 "delay must grow: {:?} -> {:?}",
2698 pair[0],
2699 pair[1]
2700 );
2701 }
2702 assert!(delays[sat..].iter().all(|d| *d == ACCEPT_BACKOFF_MAX_DELAY));
2703
2704 backoff.record_success(|| now);
2706 assert_eq!(
2707 backoff.record_exhaustion(now).delay,
2708 ACCEPT_BACKOFF_INITIAL_DELAY,
2709 );
2710
2711 let t0 = std::time::Instant::now();
2714 let mut backoff = AcceptBackoff::default();
2715 let emitted: Vec<Option<u64>> = (0..50)
2716 .map(|_| backoff.record_exhaustion(t0).log_suppressed)
2717 .collect();
2718 assert_eq!(emitted.iter().filter(|e| e.is_some()).count(), 1);
2719 assert_eq!(emitted[0], Some(0));
2720
2721 let t1 = t0 + ACCEPT_BACKOFF_LOG_INTERVAL + Duration::from_millis(1);
2722 assert_eq!(
2723 backoff.record_exhaustion(t1).log_suppressed,
2724 Some(49),
2725 "next emission reports every failure suppressed since the last one",
2726 );
2727
2728 let mut backoff = AcceptBackoff::default();
2731 let t0 = std::time::Instant::now();
2732 assert_eq!(backoff.record_exhaustion(t0).log_suppressed, Some(0));
2733 backoff.record_exhaustion(t0);
2734 backoff.record_exhaustion(t0);
2735 assert_eq!(
2736 backoff.record_success(|| t0),
2737 None,
2738 "recovery inside the log window must not emit",
2739 );
2740 backoff.record_exhaustion(t0);
2741 assert_eq!(
2742 backoff.record_success(|| t0 + ACCEPT_BACKOFF_LOG_INTERVAL),
2743 Some(3),
2744 "after the window elapses the recovery emits with the rolled-forward count",
2745 );
2746 }
2747
2748 #[cfg(unix)]
2749 #[tokio::test]
2750 async fn accept_backoff_socket_recovers_after_injected_exhaustion() {
2751 use crate::metrics::transport_metrics::TCP_ACCEPT_BACKOFF_TOTAL;
2752
2753 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
2754 .await
2755 .expect("bind ephemeral listener");
2756 let addr = listener.local_addr().expect("local_addr");
2757
2758 let mut backoff = AcceptBackoff::default();
2759 let counter_before = TCP_ACCEPT_BACKOFF_TOTAL.get();
2760
2761 let started = std::time::Instant::now();
2762 let mut expected_total = Duration::ZERO;
2763 for _ in 0..3 {
2764 expected_total += handle_accept_error(&emfile_error(), &mut backoff).await;
2765 }
2766 assert!(expected_total >= ACCEPT_BACKOFF_INITIAL_DELAY * 3);
2767 assert!(
2768 started.elapsed() >= expected_total,
2769 "the accept loop must sleep, not spin; elapsed {:?}",
2770 started.elapsed(),
2771 );
2772 assert!(TCP_ACCEPT_BACKOFF_TOTAL.get() >= counter_before + 3.0);
2773
2774 let client = tokio::spawn(async move { tokio::net::TcpStream::connect(addr).await });
2775 let (accepted, _peer) = tokio::time::timeout(Duration::from_secs(5), listener.accept())
2776 .await
2777 .expect("listener should still accept after backing off")
2778 .expect("accept should succeed");
2779 let _client = client.await.expect("client task").expect("client connect");
2780 drop(accepted);
2781
2782 backoff.record_success(std::time::Instant::now);
2783 assert_eq!(
2784 backoff.record_exhaustion(std::time::Instant::now()).delay,
2785 ACCEPT_BACKOFF_INITIAL_DELAY,
2786 );
2787 }
2788}