Skip to main content

dynamo_runtime/pipeline/network/tcp/
server.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use 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
16/// Tombstone lifetime. Bridges the `register()` → `associate_instance()`
17/// window (sub-millisecond in practice); 5s bounds the set by recent worker
18/// churn rather than process lifetime, since etcd lease IDs are unique per
19/// restart and never get cleared by an `Added` event for the same identity.
20const 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
52// Trait for IP address resolution - allows dependency injection for testing
53pub trait IpResolver {
54    fn local_ip(&self) -> Result<std::net::IpAddr, Error>;
55    fn local_ipv6(&self) -> Result<std::net::IpAddr, Error>;
56}
57
58// Default implementation using the real local_ip_address crate
59pub 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
89/// A [`TcpStreamServer`] is a TCP service that listens on a port for incoming response connections.
90/// A Response connection is a connection that is established by a client with the intention of sending
91/// specific data back to the server.
92pub struct TcpStreamServer {
93    local_ip: String,
94    local_port: u16,
95    state: Arc<Mutex<State>>,
96}
97
98// pub struct TcpStreamReceiver {
99//     address: TcpStreamConnectionInfo,
100//     state: Arc<Mutex<State>>,
101//     rx: mpsc::Receiver<ResponseType>,
102// }
103
104#[allow(dead_code)]
105struct RequestedSendConnection {
106    context: Arc<dyn AsyncEngineContext>,
107    connection: oneshot::Sender<Result<StreamSender, StreamPrologueError>>,
108    /// Capacity of the per-stream mpsc buffer between the socket task and the
109    /// engine producer; carried from the registration [`StreamOptions`].
110    send_buffer_count: usize,
111}
112
113struct RequestedRecvConnection {
114    context: Arc<dyn AsyncEngineContext>,
115    connection: oneshot::Sender<Result<StreamReceiver, StreamPrologueError>>,
116    /// Capacity of the per-stream mpsc buffer between the socket task and the
117    /// engine consumer; carried from the registration [`StreamOptions`].
118    send_buffer_count: usize,
119}
120
121/// Build the per-stream data-plane mpsc channel that bridges the socket task
122/// and the engine producer/consumer. The capacity is driven by the
123/// registration options ([`StreamOptions::send_buffer_count`]) rather than a
124/// hard-coded constant; both `process_request_stream` and
125/// `process_response_stream` size their channel through this helper. See #10293.
126fn data_plane_channel<T>(send_buffer_count: usize) -> (mpsc::Sender<T>, mpsc::Receiver<T>) {
127    // `tokio::sync::mpsc::channel` panics on a capacity of 0. Now that the value
128    // is caller-configurable via `StreamOptions::send_buffer_count`, clamp to at
129    // least 1 so a misconfigured `0` degrades to a minimal buffer instead of
130    // panicking the connection handler task.
131    mpsc::channel(send_buffer_count.max(1))
132}
133
134// /// When registering a new TcpStream on the server, the registration method will return a [`Connections`] object.
135// /// This [`Connections`] object will have two [`oneshot::Receiver`] objects, one for the [`TcpStreamSender`] and one for the [`TcpStreamReceiver`].
136// /// The [`Connections`] object can be awaited to get the [`TcpStreamSender`] and [`TcpStreamReceiver`] objects; these objects will
137// /// be made available when the matching Client has connected to the server.
138// pub struct Connections {
139//     pub address: TcpStreamConnectionInfo,
140
141//     /// The [`oneshot::Receiver`] for the [`TcpStreamSender`]. Awaiting this object will return the [`TcpStreamSender`] object once
142//     /// the client has connected to the server.
143//     pub sender: Option<oneshot::Receiver<StreamSender>>,
144
145//     /// The [`oneshot::Receiver`] for the [`TcpStreamReceiver`]. Awaiting this object will return the [`TcpStreamReceiver`] object once
146//     /// the client has connected to the server.
147//     pub receiver: Option<oneshot::Receiver<StreamReceiver>>,
148// }
149
150#[derive(Default)]
151struct State {
152    tx_subjects: HashMap<String, RequestedSendConnection>,
153    rx_subjects: HashMap<String, RequestedRecvConnection>,
154    /// subject UUID -> EndpointInstanceId. Full 4-field key isolates services
155    /// that share an endpoint name across namespaces/components.
156    subject_instance: HashMap<String, EndpointInstanceId>,
157    /// EndpointInstanceId -> tagged subject UUIDs, for batch cancellation on
158    /// removal. The `StreamType` tag tells `cancel_instance_streams` which
159    /// of `rx_subjects` / `tx_subjects` holds the registration so both halves
160    /// of a bidirectional session get dropped together.
161    instance_subjects: HashMap<EndpointInstanceId, HashSet<(StreamType, String)>>,
162    /// Tombstones (instance -> insertion time) close the
163    /// `cancel_instance_streams` vs `associate_instance` race; entries expire
164    /// after [`TOMBSTONE_TTL`].
165    removed_instances: HashMap<EndpointInstanceId, Instant>,
166    handle: Option<tokio::task::JoinHandle<Result<()>>>,
167}
168
169/// Drop tombstones older than [`TOMBSTONE_TTL`]. Called lazily on every
170/// `associate_instance` / `cancel_instance_streams` to bound the set size.
171fn 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                    // Only fall back to loopback when no routable IP exists at all;
210                    // propagate other resolver errors (I/O, platform) so
211                    // misconfigured hosts fail fast instead of silently binding
212                    // to 127.0.0.1.
213                    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        // Build TLS acceptor from environment if cert+key paths are configured.
232        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    /// Associate one or both halves of a registration with a backend instance.
252    ///
253    /// `recv_subject` is the response-stream subject (always present on TCP);
254    /// `send_subject` is the request-stream subject, set only for
255    /// bidirectional sessions. Tracking the send half here is what lets
256    /// [`Self::cancel_instance_streams`] drop the request-stream
257    /// `tx_subjects` oneshot directly when discovery removes the worker,
258    /// instead of relying on the cascade from the recv-side cancellation.
259    ///
260    /// Returns `false` if the instance is already tombstoned, in which case
261    /// both subjects are cancelled immediately and the caller should skip
262    /// `send_request` and fail with a migratable `Disconnected` error.
263    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            // Instance was already removed -- cancel immediately.
274            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    /// Cancel one pending response-stream registration. Drops the
304    /// `oneshot::Sender` so the waiting receiver resolves with `RecvError`.
305    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    /// Cancel one pending request-stream registration. Parallel to
319    /// [`Self::cancel_recv_stream`]: drops the `tx_subjects` entry and, if
320    /// the subject was associated with an instance, clears its
321    /// `(StreamType::Request, _)` tag from `instance_subjects` so the per-
322    /// instance bookkeeping stays consistent.
323    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    /// Cancel all pending streams for an instance — both response-side and
337    /// request-side halves of any bidirectional sessions tracked by
338    /// `associate_instance` — and tombstone the id so any racing associate
339    /// for the same id cancels too. Returns the number of streams cancelled.
340    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    /// Drop the tombstone for an instance that has reappeared in discovery,
365    /// so future subjects for that identity are tracked normally.
366    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    /// Build a TLS acceptor from env vars if `DYN_TCP_TLS_CERT_PATH` and
372    /// `DYN_TCP_TLS_KEY_PATH` are set. Validation and diagnostics are shared with
373    /// the request-plane server via
374    /// [`crate::tls_utils::server_tls_acceptor_config`]; a client CA enables mTLS.
375    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// todo - possible rename ResponseService to ResponseServer
455#[async_trait::async_trait]
456impl ResponseService for TcpStreamServer {
457    /// Register a new subject and sender with the response subscriber
458    /// Produces an RAII object that will deregister the subject when dropped
459    ///
460    /// we need to register both data in and data out entries
461    /// there might be forward pipeline that want to consume the data out stream
462    /// and there might be a response stream that wants to consume the data in stream
463    /// on registration, we need to specific if we want data-in, data-out or both
464    /// this will map to the type of service that is runniing, i.e. Single or Many In //
465    /// Single or Many Out
466    ///
467    /// todo(ryan) - return a connection object that can be awaited. when successfully connected,
468    /// can ask for the sender and receiver
469    ///
470    /// OR
471    ///
472    /// we make it into register sender and register receiver, both would return a connection object
473    /// and when a connection is established, we'd get the respective sender or receiver
474    ///
475    /// the registration probably needs to be done in one-go, so we should use a builder object for
476    /// requesting a receiver and optional sender
477    async fn register(&self, options: StreamOptions) -> PendingConnections {
478        // oneshot channels to pass back the sender and receiver objects
479
480        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                // Drop is sync; fire-and-forget the lock acquisition.
509                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                // Drop is sync; fire-and-forget the lock acquisition.
555                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
583/// First retry delay applied after an `AcceptFailure::Exhaustion`.
584const ACCEPT_BACKOFF_INITIAL_DELAY: Duration = Duration::from_millis(5);
585/// Ceiling the retry delay saturates at, so the listener keeps polling often
586/// enough to notice recovery while no longer spinning.
587const ACCEPT_BACKOFF_MAX_DELAY: Duration = Duration::from_secs(1);
588/// Minimum wall-clock interval between two emitted exhaustion summaries.
589const ACCEPT_BACKOFF_LOG_INTERVAL: Duration = Duration::from_secs(5);
590
591/// How the listener loop must treat a failed `accept()`.
592#[derive(Debug, Clone, Copy, PartialEq, Eq)]
593enum AcceptFailure {
594    /// The process or the host is out of file descriptors or kernel memory
595    /// — see `AcceptBackoff::classify` for the exact errno set. Retrying
596    /// immediately cannot succeed, so the loop must back off.
597    Exhaustion,
598    /// Anything else — a per-connection error that the client is expected to
599    /// retry. Keeps the historical warn-and-retry-immediately behavior.
600    Ordinary,
601}
602
603/// What the caller should do about one exhaustion failure.
604#[derive(Debug, Clone, Copy, PartialEq, Eq)]
605struct AcceptBackoffAction {
606    /// How long to sleep before calling `accept()` again.
607    delay: Duration,
608    /// `Some(n)` when the caller should emit one summary line, where `n` is the
609    /// number of failures suppressed since the previous emission. `None` means
610    /// this failure is inside the current rate-limit window and must stay silent.
611    log_suppressed: Option<u64>,
612}
613
614/// Bounded exponential retry policy for the accept loop, kept deliberately pure:
615/// it reads no clock, performs no I/O and never sleeps. The caller supplies the
616/// current `Instant` and performs the sleep, which is what lets the retry
617/// schedule and the log-rate decision be unit-tested without reproducing the
618/// host's file-descriptor limit. Vocabulary mirrors `transports::etcd::connector`'s
619/// `BackoffState`.
620#[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    /// Whether the loop is currently inside an exhaustion episode — the flag
629    /// `record_success` checks to tell a real recovery from an ordinary
630    /// steady-state accept.
631    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    /// Classify an `accept()` error. Only descriptor or kernel-memory exhaustion
650    /// gets the backoff treatment; everything else stays on the pre-existing
651    /// path so genuinely unexpected failures are not hidden or slowed down.
652    fn classify(err: &std::io::Error) -> AcceptFailure {
653        #[cfg(unix)]
654        {
655            // `std::io::ErrorKind` has no stable variant for these errnos, so
656            // the raw value is the portable-with-cfg way to recognize them.
657            // `EMFILE`/`ENFILE` are descriptor exhaustion; `ENOBUFS`/`ENOMEM`
658            // are kernel out-of-memory for the socket. All four make an
659            // immediate retry deterministic, so they get the backoff path.
660            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    /// Record one exhaustion failure and return the delay to sleep plus the
675    /// rate-limited logging decision. The delay doubles per consecutive failure
676    /// and saturates at `max_delay`.
677    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    /// Record a successful accept. `Some(suppressed)` only when this success
703    /// both ended an exhaustion episode and the shared log-rate window allows
704    /// a line — so a listener flapping at the descriptor ceiling cannot emit
705    /// an unbounded stream of recovery lines, and a suppressed recovery keeps
706    /// its count for the next emission rather than discarding it. The clock
707    /// arrives as a closure, not an `Instant`, so the steady-state accept
708    /// (every ordinary connection) hits the early return below without
709    /// reading a clock at all.
710    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
730/// Apply the accept-loop error policy to one failed `accept()`. Returns the
731/// delay that was actually slept — `Duration::ZERO` for ordinary errors, which
732/// keep the historical immediate-retry behavior.
733async fn handle_accept_error(err: &std::io::Error, backoff: &mut AcceptBackoff) -> Duration {
734    match AcceptBackoff::classify(err) {
735        AcceptFailure::Ordinary => {
736            // the client should retry, so we don't need to abort
737            tracing::warn!(error = %err, "failed to accept tcp connection");
738            // Gated like the sibling in `handle_connection`: an unconditional
739            // stderr write here doubles the log volume of any accept-error storm.
740            #[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
761// Type aliases for boxed split halves used throughout the nested handlers below.
762type BoxRead = Box<dyn tokio::io::AsyncRead + Unpin + Send>;
763type BoxWrite = Box<dyn tokio::io::AsyncWrite + Unpin + Send>;
764
765// this method listens on a tcp port for incoming connections
766// new connections are expected to send a protocol specific handshake
767// for us to determine the subject they are interested in, in this case,
768// we expect the first message to be [`FirstMessage`] from which we find
769// the sender, then we spawn a task to forward all bytes from the tcp stream
770// to the sender
771async 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        // todo - add instrumentation
804        // todo - add counter for all accepted connections
805        // todo - add gauge for all inflight connections
806        // todo - add counter for incoming bytes
807        // todo - add counter for outgoing bytes
808        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        // Spawn per-connection so the accept loop is never blocked by a slow
839        // TLS handshake. The handshake is bounded by DYN_TCP_TLS_HANDSHAKE_TIMEOUT_SECS (default 3s).
840        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    // #[instrument(level = "trace"), skip(state)]
872    // todo - clone before spawn and trace process_stream
873    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    /// This method is responsible for the internal tcp stream handshake
886    /// The handshake will specialize the stream as a request/sender or response/receiver stream
887    async fn process_stream(
888        read_half: BoxRead,
889        write_half: BoxWrite,
890        state: Arc<Mutex<State>>,
891    ) -> Result<()> {
892        // attach the codec to the reader and writer to get framed readers and writers
893        let mut framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
894        let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
895
896        // the internal tcp [`CallHomeHandshake`] connects the socket to the requester
897        // here we await this first message as a raw bytes two part message
898        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        // we await on the raw bytes which should come in as a header only message
905        // todo - improve error handling - check for no data
906        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        // branch here to handle sender stream or receiver stream
918        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    /// Symmetric to [`process_response_stream`] for the upstream→downstream
930    /// data direction: deliver the [`StreamSender`] half registered by the
931    /// upstream to whoever awaits it, then pump every frame the upstream pushes
932    /// into the now-connected TCP socket.
933    ///
934    /// One difference is that the request stream is **unidirectional**:
935    /// the upstream writes data + one closing control message, and the
936    /// downstream is not expected to reply: downstream response or inference
937    /// error should be returned through response stream. We therefore drop the
938    /// read half, on fatal error, downstream should drop the request stream.
939    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        // Request stream is unidirectional; we don't read from the downstream.
946        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        // Buffer size is driven by the registration options
962        // ([`StreamOptions::send_buffer_count`]) rather than hard-coded; the
963        // same applies to `process_response_stream`. See #10293.
964        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                // Request streams don't carry a downstream-prologue today; the
970                // upstream may begin sending immediately.
971                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    /// Pump frames the upstream queued on its `StreamSender` into the TCP socket.
985    /// The closing control message depends on why the loop exited:
986    /// - `context.killed()` → [`ControlMessage::Kill`] (hard cancel notification)
987    /// - `context.stopped()` → [`ControlMessage::Stop`] (graceful cancel notification)
988    /// - `request_rx` returns `None` → [`ControlMessage::Sentinel`] (clean EOS)
989    /// - write error → no control message, the socket is already broken
990    ///
991    /// The downstream `handle_request_reader` matches on the received variant
992    /// and reacts accordingly.
993    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        // Construct the cancellation futures once. Recreating them for every frame clones
999        // the context's watch receivers and repeatedly registers/drops Tokio notifications.
1000        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        // unwrap response_stream
1067        let RequestedRecvConnection {
1068            context,
1069            connection,
1070            send_buffer_count,
1071        } = response_stream;
1072
1073        // the [`Prologue`]
1074        // there must be a second control message it indicate the other segment's generate method was successful
1075        // No timeout here: the worker sends the prologue only after generate() setup completes,
1076        // which can take arbitrarily long (model load, queue delay, cold start).
1077        let prologue = reader
1078            .next()
1079            .await
1080            .ok_or(error!("Connection closed without a ControlMessge"))??;
1081
1082        // deserialize prologue
1083        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                        // Notify the requester as the sibling arm does. Returning on
1089                        // `?` alone drops the oneshot un-sent, and the requester then
1090                        // reports a bare disconnect that names neither the worker's
1091                        // failure nor this one.
1092                        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                // Worker sent a non-HeaderOnly frame in the prologue slot
1101                // (protocol violation, version skew, corruption). Notify the
1102                // requester so the generate call chain fails cleanly, then
1103                // return Err so the connection task ends without panicking.
1104                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        // await the control message of GTG or Error, if error, then connection.send(Err(String)), which should fail the
1111        // generate call chain
1112        //
1113        // note: this second control message might be delayed, but the expensive part of setting up the connection
1114        // is both complete and ready for data flow; awaiting here is not a performance hit or problem and it allows
1115        // us to trace the initial setup time vs the time to prologue
1116        if let Some(error) = prologue.error {
1117            let returned = error!("Received error prologue: {error}");
1118            // Forward the worker's typed error so the requesting side can classify
1119            // the failure instead of parsing the message. An older worker sends none.
1120            let _ = connection.send(Err(StreamPrologueError {
1121                message: error,
1122                typed_error: prologue.typed_error,
1123            }));
1124            return Err(returned);
1125        }
1126
1127        // Buffer size is driven by the registration options
1128        // ([`StreamOptions::send_buffer_count`]) rather than hard-coded; the
1129        // same applies to `process_request_stream`. See #10293.
1130        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        // sender task
1146        // issues control messages to the sender and when finished shuts down the socket
1147        // this should be the last task to finish and must
1148        let send_task = tokio::spawn(network_send_handler(writer, control_rx));
1149
1150        // forward task
1151        let recv_task = tokio::spawn(network_receive_handler(
1152            reader,
1153            response_tx,
1154            control_tx,
1155            context.clone(),
1156        ));
1157
1158        // check the results of each of the tasks
1159        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        // These futures stay pending across frames. Constructing them inside the loop clones
1174        // watch receivers and registers/drops notifications for every streamed token.
1175        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        // loop over reading the tcp stream and checking if the writer is closed
1181        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                    // `stopped` is now complete; keep this branch disabled because polling
1201                    // the same completed async future again would panic.
1202                    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                            // received a control message
1212                            if !header.is_empty() {
1213                                match process_control_message(header) {
1214                                    Ok(ControlAction::Continue) => {}
1215                                    Ok(ControlAction::Shutdown) => {
1216                                        if !data.is_empty() {
1217                                            // Sentinel-with-data is a protocol
1218                                            // violation; kill this stream, don't
1219                                            // assert!() the process down.
1220                                            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                                        // Malformed control message — kill only
1232                                        // this stream.
1233                                        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                            // TCP RST or decode error from worker — kill only
1249                            // this stream.
1250                            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                            // this is allowed but we try to avoid it
1256                            // the logic is that the client will tell us when its is done and the server
1257                            // will close the connection naturally when the sentinel message is received
1258                            // the client closing early represents a transport error outside the control of the
1259                            // transport library
1260                            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            // Sentinel is a worker→frontend message; receiving one here means
1279            // a producer is buggy. Skip rather than asserting — a stream-level
1280            // bug must not panic the worker.
1281            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                    // Closed enum of small variants; serialization shouldn't
1289                    // fail. If it ever does, log and skip rather than panic.
1290                    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            // the client issued a sentinel message
1322            // it has finished writing data and is now awaiting the server to close the connection
1323            tracing::trace!("sentinel received; shutting down");
1324            Ok(ControlAction::Shutdown)
1325        }
1326        ControlMessage::Kill | ControlMessage::Stop => {
1327            // Worker→frontend control direction only carries Sentinel. Kill/Stop
1328            // here is a protocol violation; the caller turns this Err into a
1329            // stream-local Kill rather than a process-fatal event.
1330            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        // Also clear the client-CA var: ambient it would turn this into an error
1366        // (client CA without a server cert/key) instead of plaintext.
1367        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        // only cert
1385        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        // only key
1393        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        // A client CA turns the response-stream server into an mTLS acceptor.
1417        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    // Mock resolver that always fails to simulate the fallback scenario
1448    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        // Test that TcpStreamServer::new works with default options
1463        // This verifies normal operation when IP detection succeeds
1464        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        // Verify the server can be used by registering a stream
1475        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        // Verify connection info is available and valid
1486        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        // Should have a valid port assigned
1497        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    /// The data-plane channel helper sizes the mpsc buffer from
1509    /// `send_buffer_count` — this is the value `process_request_stream` /
1510    /// `process_response_stream` feed it. `max_capacity()` reflects the
1511    /// channel's configured buffer, so a custom value and the default both
1512    /// reach the channel. Guards against regressing back to a hard-coded 64.
1513    #[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        // A misconfigured 0 must clamp to 1, not panic (mpsc::channel(0) panics).
1522        let (tx, _rx) = data_plane_channel::<()>(0);
1523        assert_eq!(tx.max_capacity(), 1);
1524    }
1525
1526    /// `register` must thread `StreamOptions::send_buffer_count` through to the
1527    /// stored `RequestedSendConnection` / `RequestedRecvConnection` (the
1528    /// registration structs `process_*_stream` later destructure to size the
1529    /// channel). Verified here against the real registration path.
1530    #[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        // Test fallback behavior using a mock resolver that always fails
1562        // This guarantees the fallback logic is triggered
1563
1564        let options = ServerOptions::builder().port(0).build().unwrap();
1565
1566        // Use the failing resolver to force the fallback
1567        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        // Get the actual bound address by registering a stream
1576        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        // With the failing resolver, fallback should ALWAYS be used
1596        let ip = socket_addr.ip();
1597        assert!(
1598            ip.is_loopback(),
1599            "Should use loopback when IP detection fails"
1600        );
1601
1602        // Verify it's specifically 127.0.0.1 (the fallback value from the patch)
1603        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        // The server should work with the fallback IP
1613        assert!(socket_addr.port() > 0, "Server should have a valid port");
1614    }
1615
1616    /// Create a test server using the failing IP resolver (falls back to loopback).
1617    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    /// Helper: register a response stream and extract its subject string.
1627    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    /// Convenience constructor so tests don't repeat the struct literal.
1649    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    /// Helper: register a bidirectional pair (both request + response halves)
1664    /// and return both subjects + their providers.
1665    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    /// `cancel_instance_streams` must drop the request-stream oneshot too,
1697    /// not just the response-stream one. Without the tagged tracker this test
1698    /// would hang on `send_provider.await` because the tx_subjects entry
1699    /// would leak past instance removal.
1700    #[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    /// Pre-tombstoning an instance must drop both halves of a later
1728    /// `associate_instance(recv, Some(send), id)` call, not just the recv.
1729    #[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        // Pre-tombstone the instance.
1734        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        // The oneshot receiver should now resolve with an error (sender dropped)
1763        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        // Associate first two with instance 10, third with instance 20
1779        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        // Cancel instance 10 -- should cancel 2 subjects
1784        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        // Instance 20 should be unaffected -- cancel it separately
1791        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        // Cancel the individual subject
1814        server.cancel_recv_stream(&subject).await;
1815
1816        // Instance should have no remaining subjects
1817        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        // Register a response stream but DON'T call into_parts -- just drop it
1829        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        // Get the subject before dropping
1841        let tcp_info: TcpStreamConnectionInfo =
1842            recv_stream.connection_info.clone().try_into().unwrap();
1843        let subject = tcp_info.subject.clone();
1844
1845        // Verify it's in rx_subjects
1846        {
1847            let state = server.state.lock();
1848            assert!(state.rx_subjects.contains_key(&subject));
1849        }
1850
1851        // Drop the RegisteredStream -- RAII cleanup should fire
1852        drop(recv_stream);
1853
1854        // Give the spawned cleanup task a moment to run
1855        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1856
1857        // Verify it's been removed from rx_subjects
1858        {
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        // Call into_parts to disarm the cleanup
1887        let (_conn_info, _provider) = recv_stream.into_parts();
1888
1889        // Give any potential cleanup a moment to run
1890        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1891
1892        // The entry should still be in rx_subjects (cleanup was disarmed)
1893        {
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        // Simulates the race: cancel_instance_streams fires before associate_instance.
1905        let server = test_server().await;
1906
1907        let id = make_eid("ns", "comp", "generate", 42);
1908
1909        // Cancel BEFORE any subject is registered (tombstone).
1910        let cancelled = server.cancel_instance_streams(&id).await;
1911        assert_eq!(cancelled, 0);
1912
1913        // Now register a subject and try to associate it with the tombstoned instance.
1914        let (subject, provider) = register_and_get_subject(&server).await;
1915        let associated = server.associate_instance(&subject, None, &id).await;
1916
1917        // associate_instance should return false when the instance is tombstoned.
1918        assert!(
1919            !associated,
1920            "associate_instance on a tombstoned instance should return false"
1921        );
1922
1923        // The provider should resolve with an error because associate_instance
1924        // found the tombstone and immediately cancelled the subject.
1925        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        // Now associate should work normally (subject NOT cancelled).
1942        let (subject, _provider) = register_and_get_subject(&server).await;
1943        assert!(server.associate_instance(&subject, None, &id).await);
1944
1945        // Subject should be tracked, not cancelled.
1946        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        // Regression: cancelling "generate" must not cancel "prefill" subjects
1956        // that share the same instance_id (same backend runtime).
1957        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        // Cancel only the "generate" endpoint's subjects.
1969        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        // prefill must still be tracked.
1977        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        // Tombstoning "generate" must not prevent new associations on "prefill"
1985        // for the same instance_id.
1986        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        // A new subject for "generate" should be rejected.
1994        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        // A new subject for "prefill" with the same instance_id should be accepted.
2002        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        // Regression: two services with different (namespace, component) but the
2014        // same endpoint name and the same pod-backed instance_id must not interfere,
2015        // even though they share a single TcpStreamServer runtime.
2016        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        // Same endpoint name + instance_id, different namespace/component.
2022        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        // Cancel service A -- only subj_a should be affected.
2029        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        // Service B subject must still be pending.
2034        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        // After TOMBSTONE_TTL elapses, a previously-tombstoned identity must
2042        // accept new associations again, AND the entry must be physically
2043        // pruned from `removed_instances` so the set remains bounded.
2044        let server = test_server().await;
2045
2046        let id = make_eid("ns", "comp", "generate", 42);
2047
2048        // Tombstone the identity.
2049        server.cancel_instance_streams(&id).await;
2050        {
2051            let state = server.state.lock();
2052            assert!(state.removed_instances.contains_key(&id));
2053        }
2054
2055        // Advance past the TTL.
2056        tokio::time::advance(TOMBSTONE_TTL + Duration::from_secs(1)).await;
2057
2058        // associate_instance for the same identity should now succeed (no
2059        // longer tombstoned). Any new subject must be tracked normally.
2060        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        // The expired tombstone must have been pruned (lazy pruning fires on
2067        // every associate_instance/cancel_instance_streams call).
2068        {
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        // Regression net for the original tombstone fix: a tombstone younger
2080        // than TTL must still cancel late-arriving associate_instance() calls.
2081        let server = test_server().await;
2082
2083        let id = make_eid("ns", "comp", "generate", 42);
2084        server.cancel_instance_streams(&id).await;
2085
2086        // Advance only a small fraction of the TTL.
2087        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        // Old tombstones must be pruned on the next cancel_instance_streams
2100        // call, regardless of which identity is being tombstoned.
2101        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        // Documents the monotonic-lease invariant: `clear_instance_tombstone`
2125        // for one EndpointInstanceId must not touch a sibling entry. With etcd
2126        // lease IDs this defensive code rarely fires (new lease = new
2127        // EndpointInstanceId), but the per-key scope must hold.
2128        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        // A tombstone on (ns-a, comp-a, generate, 42) must not block
2146        // associations on (ns-b, comp-b, generate, 42).
2147        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        // Tombstone only service A.
2153        server.cancel_instance_streams(&id_a).await;
2154
2155        // Service A is tombstoned — new association is rejected.
2156        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        // Service B with same endpoint name + instance_id must be accepted.
2161        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    /// Stand up a TcpStreamServer, register a response stream, connect a
2174    /// client, drive the handshake + prologue, and return the client-side
2175    /// framed reader/writer along with the receiver.
2176    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        // SAFETY (test-only): healthy localhost handshake always resolves all
2221        // three layers; a panic here means the harness is broken.
2222        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        // SAFETY (test-only): a misbehaving server in any of these layers is
2233        // exactly the harness failure we want surfaced as a test panic.
2234        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    /// Sending an unexpected control message (Stop or Kill from the data
2245    /// direction) is a protocol violation. The server's
2246    /// network_receive_handler must reply with ControlMessage::Kill on
2247    /// that stream alone, not panic.
2248    #[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    /// A framing/decode error from the worker side is unrecoverable for
2268    /// this stream but must not panic the worker. Server should send Kill
2269    /// and tear down only this connection.
2270    #[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    /// Sentinel is supposed to be header-only. A misbehaving client that
2286    /// attaches a data payload must not panic the worker via assert!().
2287    #[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    /// The prologue must be a HeaderOnly frame. A non-HeaderOnly prologue
2311    /// (data-only or mixed) must surface as Err to the requester rather
2312    /// than panic the worker.
2313    #[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        // Send a data-only frame in the prologue slot.
2347        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        // StreamReceiver doesn't impl Debug, so we can't use `.expect_err`.
2359        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    /// A prologue from a newer worker may carry an `ErrorType` this build does
2369    /// not know; its legacy error must still reach the requester.
2370    #[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        // Correctly framed HeaderOnly, but the typed error names a variant this
2404        // build does not know.
2405        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        // `StreamReceiver` is not `Debug`, so match instead of `expect_err`.
2418        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    /// A worker that refuses a request before producing any response bytes must
2485    /// keep its error type all the way to the requesting side.
2486    #[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        // `StreamReceiver` is not `Debug`, so match instead of `expect_err`.
2525        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    // ==================== request_stream_send_handler integration tests ====================
2537    //
2538    // These exercise the closing-message contract of `request_stream_send_handler`
2539    // end-to-end: register a request stream, dial it as a raw client (so we can
2540    // inspect frames directly), then trigger each of the exit branches and
2541    // assert which ControlMessage arrives on the wire.
2542
2543    use futures::SinkExt;
2544
2545    /// Register a request stream and dial it with a raw client. Returns the
2546    /// framed reader on the raw client side, the StreamSender held by the
2547    /// upstream, and the upstream's engine context (so the test can drive
2548    /// kill / stop externally).
2549    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    /// Pull frames off the raw client reader until the first `ControlMessage`
2590    /// arrives, ignoring any DataOnly frames before it. Returns the variant.
2591    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            // DataOnly frame — skip and keep reading.
2605        }
2606    }
2607
2608    /// Dropping the upstream's StreamSender drains `request_rx` and the server
2609    /// emits [`ControlMessage::Sentinel`] as the closing frame.
2610    #[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    /// `context.kill()` makes the server emit [`ControlMessage::Kill`] before
2625    /// shutting down the write half.
2626    #[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    /// `context.stop()` makes the server emit [`ControlMessage::Stop`] before
2641    /// shutting down the write half.
2642    #[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    // ---- accept-loop backoff under resource exhaustion (issue #11822) ----
2657
2658    #[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        // Ordinary errors stay on the immediate-retry path.
2675        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        // Delay doubles per consecutive failure and saturates at the ceiling.
2686        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        // A successful accept resets the schedule to the initial delay.
2705        backoff.record_success(|| now);
2706        assert_eq!(
2707            backoff.record_exhaustion(now).delay,
2708            ACCEPT_BACKOFF_INITIAL_DELAY,
2709        );
2710
2711        // The summary is rate-limited: one emission per interval, carrying the
2712        // count of failures suppressed since the previous one.
2713        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        // Recovery shares the rate-limit window with the failure summary, so a
2729        // flapping listener cannot emit a recovery line per accept.
2730        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}