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, 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::sync::Mutex;
13use tokio::time::Instant;
14
15/// Tombstone lifetime. Bridges the `register()` → `associate_instance()`
16/// window (sub-millisecond in practice); 5s bounds the set by recent worker
17/// churn rather than process lifetime, since etcd lease IDs are unique per
18/// restart and never get cleared by an `Added` event for the same identity.
19const TOMBSTONE_TTL: Duration = Duration::from_secs(5);
20
21use bytes::Bytes;
22use derive_builder::Builder;
23use futures::{SinkExt, StreamExt};
24use local_ip_address::{Error, list_afinet_netifas, local_ip, local_ipv6};
25
26use serde::{Deserialize, Serialize};
27use tokio::{
28    io::AsyncWriteExt,
29    sync::{mpsc, oneshot},
30    time,
31};
32use tokio_util::codec::{FramedRead, FramedWrite};
33
34use super::{
35    CallHomeHandshake, ControlMessage, PendingConnections, RegisteredStream, StreamOptions,
36    StreamReceiver, StreamSender, TcpStreamConnectionInfo, TwoPartCodec,
37};
38use crate::discovery::EndpointInstanceId;
39use crate::engine::AsyncEngineContext;
40use crate::pipeline::{
41    PipelineError,
42    network::{
43        ResponseService, ResponseStreamPrologue,
44        codec::{TwoPartMessage, TwoPartMessageType},
45        tcp::StreamType,
46    },
47};
48use anyhow::{Context, Result, anyhow as error};
49
50// Trait for IP address resolution - allows dependency injection for testing
51pub trait IpResolver {
52    fn local_ip(&self) -> Result<std::net::IpAddr, Error>;
53    fn local_ipv6(&self) -> Result<std::net::IpAddr, Error>;
54}
55
56// Default implementation using the real local_ip_address crate
57pub struct DefaultIpResolver;
58
59impl IpResolver for DefaultIpResolver {
60    fn local_ip(&self) -> Result<std::net::IpAddr, Error> {
61        local_ip()
62    }
63
64    fn local_ipv6(&self) -> Result<std::net::IpAddr, Error> {
65        local_ipv6()
66    }
67}
68
69#[allow(dead_code)]
70type ResponseType = TwoPartMessage;
71
72#[derive(Debug, Serialize, Deserialize, Clone, Builder, Default)]
73pub struct ServerOptions {
74    #[builder(default = "0")]
75    pub port: u16,
76
77    #[builder(default)]
78    pub interface: Option<String>,
79}
80
81impl ServerOptions {
82    pub fn builder() -> ServerOptionsBuilder {
83        ServerOptionsBuilder::default()
84    }
85}
86
87/// A [`TcpStreamServer`] is a TCP service that listens on a port for incoming response connections.
88/// A Response connection is a connection that is established by a client with the intention of sending
89/// specific data back to the server.
90pub struct TcpStreamServer {
91    local_ip: String,
92    local_port: u16,
93    state: Arc<Mutex<State>>,
94}
95
96// pub struct TcpStreamReceiver {
97//     address: TcpStreamConnectionInfo,
98//     state: Arc<Mutex<State>>,
99//     rx: mpsc::Receiver<ResponseType>,
100// }
101
102#[allow(dead_code)]
103struct RequestedSendConnection {
104    context: Arc<dyn AsyncEngineContext>,
105    connection: oneshot::Sender<Result<StreamSender, String>>,
106    /// Capacity of the per-stream mpsc buffer between the socket task and the
107    /// engine producer; carried from the registration [`StreamOptions`].
108    send_buffer_count: usize,
109}
110
111struct RequestedRecvConnection {
112    context: Arc<dyn AsyncEngineContext>,
113    connection: oneshot::Sender<Result<StreamReceiver, String>>,
114    /// Capacity of the per-stream mpsc buffer between the socket task and the
115    /// engine consumer; carried from the registration [`StreamOptions`].
116    send_buffer_count: usize,
117}
118
119/// Build the per-stream data-plane mpsc channel that bridges the socket task
120/// and the engine producer/consumer. The capacity is driven by the
121/// registration options ([`StreamOptions::send_buffer_count`]) rather than a
122/// hard-coded constant; both `process_request_stream` and
123/// `process_response_stream` size their channel through this helper. See #10293.
124fn data_plane_channel<T>(send_buffer_count: usize) -> (mpsc::Sender<T>, mpsc::Receiver<T>) {
125    // `tokio::sync::mpsc::channel` panics on a capacity of 0. Now that the value
126    // is caller-configurable via `StreamOptions::send_buffer_count`, clamp to at
127    // least 1 so a misconfigured `0` degrades to a minimal buffer instead of
128    // panicking the connection handler task.
129    mpsc::channel(send_buffer_count.max(1))
130}
131
132// /// When registering a new TcpStream on the server, the registration method will return a [`Connections`] object.
133// /// This [`Connections`] object will have two [`oneshot::Receiver`] objects, one for the [`TcpStreamSender`] and one for the [`TcpStreamReceiver`].
134// /// The [`Connections`] object can be awaited to get the [`TcpStreamSender`] and [`TcpStreamReceiver`] objects; these objects will
135// /// be made available when the matching Client has connected to the server.
136// pub struct Connections {
137//     pub address: TcpStreamConnectionInfo,
138
139//     /// The [`oneshot::Receiver`] for the [`TcpStreamSender`]. Awaiting this object will return the [`TcpStreamSender`] object once
140//     /// the client has connected to the server.
141//     pub sender: Option<oneshot::Receiver<StreamSender>>,
142
143//     /// The [`oneshot::Receiver`] for the [`TcpStreamReceiver`]. Awaiting this object will return the [`TcpStreamReceiver`] object once
144//     /// the client has connected to the server.
145//     pub receiver: Option<oneshot::Receiver<StreamReceiver>>,
146// }
147
148#[derive(Default)]
149struct State {
150    tx_subjects: HashMap<String, RequestedSendConnection>,
151    rx_subjects: HashMap<String, RequestedRecvConnection>,
152    /// subject UUID -> EndpointInstanceId. Full 4-field key isolates services
153    /// that share an endpoint name across namespaces/components.
154    subject_instance: HashMap<String, EndpointInstanceId>,
155    /// EndpointInstanceId -> tagged subject UUIDs, for batch cancellation on
156    /// removal. The `StreamType` tag tells `cancel_instance_streams` which
157    /// of `rx_subjects` / `tx_subjects` holds the registration so both halves
158    /// of a bidirectional session get dropped together.
159    instance_subjects: HashMap<EndpointInstanceId, HashSet<(StreamType, String)>>,
160    /// Tombstones (instance -> insertion time) close the
161    /// `cancel_instance_streams` vs `associate_instance` race; entries expire
162    /// after [`TOMBSTONE_TTL`].
163    removed_instances: HashMap<EndpointInstanceId, Instant>,
164    handle: Option<tokio::task::JoinHandle<Result<()>>>,
165}
166
167/// Drop tombstones older than [`TOMBSTONE_TTL`]. Called lazily on every
168/// `associate_instance` / `cancel_instance_streams` to bound the set size.
169fn prune_tombstones(tombstones: &mut HashMap<EndpointInstanceId, Instant>, now: Instant) {
170    tombstones.retain(|_, ts| now.saturating_duration_since(*ts) < TOMBSTONE_TTL);
171}
172
173impl TcpStreamServer {
174    pub fn options_builder() -> ServerOptionsBuilder {
175        ServerOptionsBuilder::default()
176    }
177
178    pub async fn new(options: ServerOptions) -> Result<Arc<Self>, PipelineError> {
179        Self::new_with_resolver(options, DefaultIpResolver).await
180    }
181
182    pub async fn new_with_resolver<R: IpResolver>(
183        options: ServerOptions,
184        resolver: R,
185    ) -> Result<Arc<Self>, PipelineError> {
186        let local_ip = match options.interface {
187            Some(interface) => {
188                let interfaces: HashMap<String, std::net::IpAddr> =
189                    list_afinet_netifas()?.into_iter().collect();
190
191                interfaces
192                    .get(&interface)
193                    .ok_or(PipelineError::Generic(format!(
194                        "Interface not found: {}",
195                        interface
196                    )))?
197                    .to_string()
198            }
199            None => {
200                let resolved_ip = resolver.local_ip().or_else(|err| match err {
201                    Error::LocalIpAddressNotFound => resolver.local_ipv6(),
202                    _ => Err(err),
203                });
204
205                match resolved_ip {
206                    Ok(addr) => addr,
207                    // Only fall back to loopback when no routable IP exists at all;
208                    // propagate other resolver errors (I/O, platform) so
209                    // misconfigured hosts fail fast instead of silently binding
210                    // to 127.0.0.1.
211                    Err(Error::LocalIpAddressNotFound) => {
212                        tracing::warn!(
213                            "No routable local IP address found; falling back to 127.0.0.1"
214                        );
215                        IpAddr::from([127, 0, 0, 1])
216                    }
217                    Err(err) => {
218                        return Err(PipelineError::Generic(format!(
219                            "Failed to resolve local IP address: {err}"
220                        )));
221                    }
222                }
223                .to_string()
224            }
225        };
226
227        let state = Arc::new(Mutex::new(State::default()));
228
229        let local_port = Self::start(local_ip.clone(), options.port, state.clone())
230            .await
231            .map_err(|e| {
232                PipelineError::Generic(format!("Failed to start TcpStreamServer: {}", e))
233            })?;
234
235        tracing::debug!("tcp transport service on {local_ip}:{local_port}");
236
237        Ok(Arc::new(Self {
238            local_ip,
239            local_port,
240            state,
241        }))
242    }
243
244    /// Associate one or both halves of a registration with a backend instance.
245    ///
246    /// `recv_subject` is the response-stream subject (always present on TCP);
247    /// `send_subject` is the request-stream subject, set only for
248    /// bidirectional sessions. Tracking the send half here is what lets
249    /// [`Self::cancel_instance_streams`] drop the request-stream
250    /// `tx_subjects` oneshot directly when discovery removes the worker,
251    /// instead of relying on the cascade from the recv-side cancellation.
252    ///
253    /// Returns `false` if the instance is already tombstoned, in which case
254    /// both subjects are cancelled immediately and the caller should skip
255    /// `send_request` and fail with a migratable `Disconnected` error.
256    pub async fn associate_instance(
257        &self,
258        recv_subject: &str,
259        send_subject: Option<&str>,
260        id: &EndpointInstanceId,
261    ) -> bool {
262        let mut state = self.state.lock().await;
263        let now = Instant::now();
264        prune_tombstones(&mut state.removed_instances, now);
265        if state.removed_instances.contains_key(id) {
266            // Instance was already removed -- cancel immediately.
267            tracing::warn!(
268                recv_subject,
269                send_subject,
270                namespace = %id.namespace,
271                component = %id.component,
272                endpoint = %id.endpoint,
273                instance_id = id.instance_id,
274                "Cancelling subject immediately: instance already removed (tombstoned)"
275            );
276            state.rx_subjects.remove(recv_subject);
277            if let Some(s) = send_subject {
278                state.tx_subjects.remove(s);
279            }
280            return false;
281        }
282        state
283            .subject_instance
284            .insert(recv_subject.to_string(), id.clone());
285        if let Some(s) = send_subject {
286            state.subject_instance.insert(s.to_string(), id.clone());
287        }
288        let entry = state.instance_subjects.entry(id.clone()).or_default();
289        entry.insert((StreamType::Response, recv_subject.to_string()));
290        if let Some(s) = send_subject {
291            entry.insert((StreamType::Request, s.to_string()));
292        }
293        true
294    }
295
296    /// Cancel one pending response-stream registration. Drops the
297    /// `oneshot::Sender` so the waiting receiver resolves with `RecvError`.
298    pub async fn cancel_recv_stream(&self, subject: &str) {
299        let mut state = self.state.lock().await;
300        state.rx_subjects.remove(subject);
301        if let Some(key) = state.subject_instance.remove(subject)
302            && let Some(subjects) = state.instance_subjects.get_mut(&key)
303        {
304            subjects.remove(&(StreamType::Response, subject.to_string()));
305            if subjects.is_empty() {
306                state.instance_subjects.remove(&key);
307            }
308        }
309    }
310
311    /// Cancel one pending request-stream registration. Parallel to
312    /// [`Self::cancel_recv_stream`]: drops the `tx_subjects` entry and, if
313    /// the subject was associated with an instance, clears its
314    /// `(StreamType::Request, _)` tag from `instance_subjects` so the per-
315    /// instance bookkeeping stays consistent.
316    pub async fn cancel_send_stream(&self, subject: &str) {
317        let mut state = self.state.lock().await;
318        state.tx_subjects.remove(subject);
319        if let Some(key) = state.subject_instance.remove(subject)
320            && let Some(subjects) = state.instance_subjects.get_mut(&key)
321        {
322            subjects.remove(&(StreamType::Request, subject.to_string()));
323            if subjects.is_empty() {
324                state.instance_subjects.remove(&key);
325            }
326        }
327    }
328
329    /// Cancel all pending streams for an instance — both response-side and
330    /// request-side halves of any bidirectional sessions tracked by
331    /// `associate_instance` — and tombstone the id so any racing associate
332    /// for the same id cancels too. Returns the number of streams cancelled.
333    pub async fn cancel_instance_streams(&self, id: &EndpointInstanceId) -> usize {
334        let mut state = self.state.lock().await;
335        let now = Instant::now();
336        prune_tombstones(&mut state.removed_instances, now);
337        state.removed_instances.insert(id.clone(), now);
338        let subjects = match state.instance_subjects.remove(id) {
339            Some(subjects) => subjects,
340            None => return 0,
341        };
342        let count = subjects.len();
343        for (kind, subject) in &subjects {
344            match kind {
345                StreamType::Response => {
346                    state.rx_subjects.remove(subject);
347                }
348                StreamType::Request => {
349                    state.tx_subjects.remove(subject);
350                }
351            }
352            state.subject_instance.remove(subject);
353        }
354        count
355    }
356
357    /// Drop the tombstone for an instance that has reappeared in discovery,
358    /// so future subjects for that identity are tracked normally.
359    pub async fn clear_instance_tombstone(&self, id: &EndpointInstanceId) {
360        let mut state = self.state.lock().await;
361        state.removed_instances.remove(id);
362    }
363
364    #[allow(clippy::await_holding_lock)]
365    async fn start(local_ip: String, local_port: u16, state: Arc<Mutex<State>>) -> Result<u16> {
366        let addr = format!("{}:{}", local_ip, local_port);
367        let state_clone = state.clone();
368        let mut guard = state.lock().await;
369        if guard.handle.is_some() {
370            panic!("TcpStreamServer already started");
371        }
372        let (ready_tx, ready_rx) = tokio::sync::oneshot::channel::<Result<u16>>();
373        let handle = tokio::spawn(tcp_listener(addr, state_clone, ready_tx));
374        guard.handle = Some(handle);
375        drop(guard);
376        let local_port = ready_rx.await??;
377        Ok(local_port)
378    }
379}
380
381// todo - possible rename ResponseService to ResponseServer
382#[async_trait::async_trait]
383impl ResponseService for TcpStreamServer {
384    /// Register a new subject and sender with the response subscriber
385    /// Produces an RAII object that will deregister the subject when dropped
386    ///
387    /// we need to register both data in and data out entries
388    /// there might be forward pipeline that want to consume the data out stream
389    /// and there might be a response stream that wants to consume the data in stream
390    /// on registration, we need to specific if we want data-in, data-out or both
391    /// this will map to the type of service that is runniing, i.e. Single or Many In //
392    /// Single or Many Out
393    ///
394    /// todo(ryan) - return a connection object that can be awaited. when successfully connected,
395    /// can ask for the sender and receiver
396    ///
397    /// OR
398    ///
399    /// we make it into register sender and register receiver, both would return a connection object
400    /// and when a connection is established, we'd get the respective sender or receiver
401    ///
402    /// the registration probably needs to be done in one-go, so we should use a builder object for
403    /// requesting a receiver and optional sender
404    async fn register(&self, options: StreamOptions) -> PendingConnections {
405        // oneshot channels to pass back the sender and receiver objects
406
407        let address = format!("{}:{}", self.local_ip, self.local_port);
408        tracing::debug!("Registering new TcpStream on {address}");
409
410        let send_stream = if options.enable_request_stream {
411            let sender_subject = uuid::Uuid::new_v4().to_string();
412
413            let (pending_sender_tx, pending_sender_rx) = oneshot::channel();
414
415            let connection_info = RequestedSendConnection {
416                context: options.context.clone(),
417                connection: pending_sender_tx,
418                send_buffer_count: options.send_buffer_count,
419            };
420
421            let mut state = self.state.lock().await;
422            state
423                .tx_subjects
424                .insert(sender_subject.clone(), connection_info);
425
426            let cleanup_subject = sender_subject.clone();
427            let cleanup_state = self.state.clone();
428            let registered_stream = RegisteredStream::new(
429                TcpStreamConnectionInfo {
430                    address: address.clone(),
431                    subject: sender_subject,
432                    context: options.context.id().to_string(),
433                    stream_type: StreamType::Request,
434                }
435                .into(),
436                pending_sender_rx,
437            )
438            .with_cleanup(move || {
439                // Drop is sync; fire-and-forget the lock acquisition.
440                tokio::spawn(async move {
441                    let mut state = cleanup_state.lock().await;
442                    state.tx_subjects.remove(&cleanup_subject);
443                    if let Some(key) = state.subject_instance.remove(&cleanup_subject)
444                        && let Some(subjects) = state.instance_subjects.get_mut(&key)
445                    {
446                        subjects.remove(&(StreamType::Request, cleanup_subject.clone()));
447                        if subjects.is_empty() {
448                            state.instance_subjects.remove(&key);
449                        }
450                    }
451                });
452            });
453
454            Some(registered_stream)
455        } else {
456            None
457        };
458
459        let recv_stream = if options.enable_response_stream {
460            let (pending_recver_tx, pending_recver_rx) = oneshot::channel();
461            let receiver_subject = uuid::Uuid::new_v4().to_string();
462
463            let connection_info = RequestedRecvConnection {
464                context: options.context.clone(),
465                connection: pending_recver_tx,
466                send_buffer_count: options.send_buffer_count,
467            };
468
469            let mut state = self.state.lock().await;
470            state
471                .rx_subjects
472                .insert(receiver_subject.clone(), connection_info);
473
474            let cleanup_subject = receiver_subject.clone();
475            let cleanup_state = self.state.clone();
476            let registered_stream = RegisteredStream::new(
477                TcpStreamConnectionInfo {
478                    address: address.clone(),
479                    subject: receiver_subject,
480                    context: options.context.id().to_string(),
481                    stream_type: StreamType::Response,
482                }
483                .into(),
484                pending_recver_rx,
485            )
486            .with_cleanup(move || {
487                // Drop is sync; fire-and-forget the lock acquisition.
488                tokio::spawn(async move {
489                    let mut state = cleanup_state.lock().await;
490                    state.rx_subjects.remove(&cleanup_subject);
491                    if let Some(key) = state.subject_instance.remove(&cleanup_subject)
492                        && let Some(subjects) = state.instance_subjects.get_mut(&key)
493                    {
494                        subjects.remove(&(StreamType::Response, cleanup_subject.clone()));
495                        if subjects.is_empty() {
496                            state.instance_subjects.remove(&key);
497                        }
498                    }
499                });
500            });
501
502            Some(registered_stream)
503        } else {
504            None
505        };
506
507        PendingConnections {
508            send_stream,
509            recv_stream,
510        }
511    }
512}
513
514// this method listens on a tcp port for incoming connections
515// new connections are expected to send a protocol specific handshake
516// for us to determine the subject they are interested in, in this case,
517// we expect the first message to be [`FirstMessage`] from which we find
518// the sender, then we spawn a task to forward all bytes from the tcp stream
519// to the sender
520async fn tcp_listener(
521    addr: String,
522    state: Arc<Mutex<State>>,
523    read_tx: tokio::sync::oneshot::Sender<Result<u16>>,
524) -> Result<()> {
525    let listener = tokio::net::TcpListener::bind(&addr)
526        .await
527        .map_err(|e| anyhow::anyhow!("Failed to start TcpListender on {}: {}", addr, e));
528
529    let listener = match listener {
530        Ok(listener) => {
531            let addr = listener
532                .local_addr()
533                .map_err(|e| anyhow::anyhow!("Failed get SocketAddr: {:?}", e))
534                .unwrap();
535
536            read_tx
537                .send(Ok(addr.port()))
538                .expect("Failed to send ready signal");
539
540            listener
541        }
542        Err(e) => {
543            read_tx.send(Err(e)).expect("Failed to send ready signal");
544            return Err(anyhow::anyhow!("Failed to start TcpListender on {}", addr));
545        }
546    };
547
548    loop {
549        // todo - add instrumentation
550        // todo - add counter for all accepted connections
551        // todo - add gauge for all inflight connections
552        // todo - add counter for incoming bytes
553        // todo - add counter for outgoing bytes
554        let (stream, _addr) = match listener.accept().await {
555            Ok((stream, _addr)) => (stream, _addr),
556            Err(e) => {
557                // the client should retry, so we don't need to abort
558                tracing::warn!("failed to accept tcp connection: {e}");
559                eprintln!("failed to accept tcp connection: {}", e);
560                continue;
561            }
562        };
563
564        match stream.set_nodelay(true) {
565            Ok(_) => (),
566            Err(e) => {
567                tracing::warn!("failed to set tcp stream to nodelay: {e}");
568            }
569        }
570
571        match stream.set_linger(Some(std::time::Duration::from_secs(0))) {
572            Ok(_) => (),
573            Err(e) => {
574                tracing::warn!("failed to set tcp stream to linger: {e}");
575            }
576        }
577
578        tokio::spawn(handle_connection(stream, state.clone()));
579    }
580
581    // #[instrument(level = "trace"), skip(state)]
582    // todo - clone before spawn and trace process_stream
583    async fn handle_connection(stream: tokio::net::TcpStream, state: Arc<Mutex<State>>) {
584        let result = process_stream(stream, state).await;
585        match result {
586            Ok(_) => tracing::trace!("successfully processed tcp connection"),
587            Err(e) => {
588                tracing::warn!("failed to handle tcp connection: {e}");
589                #[cfg(debug_assertions)]
590                eprintln!("failed to handle tcp connection: {}", e);
591            }
592        }
593    }
594
595    /// This method is responsible for the internal tcp stream handshake
596    /// The handshake will specialize the stream as a request/sender or response/receiver stream
597    async fn process_stream(stream: tokio::net::TcpStream, state: Arc<Mutex<State>>) -> Result<()> {
598        // split the socket in to a reader and writer
599        let (read_half, write_half) = tokio::io::split(stream);
600
601        // attach the codec to the reader and writer to get framed readers and writers
602        let mut framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
603        let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
604
605        // the internal tcp [`CallHomeHandshake`] connects the socket to the requester
606        // here we await this first message as a raw bytes two part message
607        let first_message = framed_reader
608            .next()
609            .await
610            .ok_or(error!("Connection closed without a ControlMessage"))??;
611
612        // we await on the raw bytes which should come in as a header only message
613        // todo - improve error handling - check for no data
614        let handshake: CallHomeHandshake = match first_message.header() {
615            Some(header) => serde_json::from_slice(header).map_err(|e| {
616                error!(
617                    "Failed to deserialize the first message as a valid `CallHomeHandshake`: {e}",
618                )
619            })?,
620            None => {
621                return Err(error!("Expected ControlMessage, got DataMessage"));
622            }
623        };
624
625        // branch here to handle sender stream or receiver stream
626        match handshake.stream_type {
627            StreamType::Request => {
628                process_request_stream(handshake.subject, state, framed_reader, framed_writer).await
629            }
630            StreamType::Response => {
631                process_response_stream(handshake.subject, state, framed_reader, framed_writer)
632                    .await
633            }
634        }
635    }
636
637    /// Symmetric to [`process_response_stream`] for the upstream→downstream
638    /// data direction: deliver the [`StreamSender`] half registered by the
639    /// upstream to whoever awaits it, then pump every frame the upstream pushes
640    /// into the now-connected TCP socket.
641    ///
642    /// One difference is that the request stream is **unidirectional**:
643    /// the upstream writes data + one closing control message, and the
644    /// downstream is not expected to reply: downstream response or inference
645    /// error should be returned through response stream. We therefore drop the
646    /// read half, on fatal error, downstream should drop the request stream.
647    async fn process_request_stream(
648        subject: String,
649        state: Arc<Mutex<State>>,
650        reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
651        writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
652    ) -> Result<()> {
653        // Request stream is unidirectional; we don't read from the downstream.
654        drop(reader);
655
656        let request_stream = {
657            let mut guard = state.lock().await;
658            let conn = guard.tx_subjects.remove(&subject).ok_or(error!(
659                "Subject not found: {}; downstream subscriber specified a subject unknown to the upstream publisher",
660                subject
661            ))?;
662            if let Some(key) = guard.subject_instance.remove(&subject)
663                && let Some(subjects) = guard.instance_subjects.get_mut(&key)
664            {
665                subjects.remove(&(StreamType::Request, subject.clone()));
666                if subjects.is_empty() {
667                    guard.instance_subjects.remove(&key);
668                }
669            }
670            conn
671        };
672
673        let RequestedSendConnection {
674            context,
675            connection,
676            send_buffer_count,
677        } = request_stream;
678
679        // Buffer size is driven by the registration options
680        // ([`StreamOptions::send_buffer_count`]) rather than hard-coded; the
681        // same applies to `process_response_stream`. See #10293.
682        let (request_tx, request_rx) = data_plane_channel(send_buffer_count);
683
684        if connection
685            .send(Ok(crate::pipeline::network::StreamSender {
686                tx: request_tx,
687                // Request streams don't carry a downstream-prologue today; the
688                // upstream may begin sending immediately.
689                prologue: None,
690            }))
691            .is_err()
692        {
693            return Err(error!(
694                "The requester of the request stream has been dropped before the connection was established"
695            ));
696        }
697
698        request_stream_send_handler(writer, request_rx, context).await;
699        Ok(())
700    }
701
702    /// Pump frames the upstream queued on its `StreamSender` into the TCP socket.
703    /// The closing control message depends on why the loop exited:
704    /// - `context.killed()` → [`ControlMessage::Kill`] (hard cancel notification)
705    /// - `context.stopped()` → [`ControlMessage::Stop`] (graceful cancel notification)
706    /// - `request_rx` returns `None` → [`ControlMessage::Sentinel`] (clean EOS)
707    /// - write error → no control message, the socket is already broken
708    ///
709    /// The downstream `handle_request_reader` matches on the received variant
710    /// and reacts accordingly.
711    async fn request_stream_send_handler(
712        mut framed_writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
713        mut request_rx: mpsc::Receiver<TwoPartMessage>,
714        context: Arc<dyn AsyncEngineContext>,
715    ) {
716        let closing_msg: Option<ControlMessage> = loop {
717            tokio::select! {
718                biased;
719
720                _ = context.killed() => {
721                    tracing::trace!("context kill received in request-stream send handler");
722                    break Some(ControlMessage::Kill);
723                }
724
725                _ = context.stopped() => {
726                    tracing::trace!("context stop received in request-stream send handler");
727                    break Some(ControlMessage::Stop);
728                }
729
730                msg = request_rx.recv() => {
731                    match msg {
732                        Some(msg) => {
733                            if let Err(e) = framed_writer.send(msg).await {
734                                tracing::trace!(
735                                    "failed to send request-stream frame to downstream: {:?}",
736                                    e
737                                );
738                                break None;
739                            }
740                        }
741                        None => {
742                            tracing::trace!("upstream request-stream sender closed; sending sentinel");
743                            break Some(ControlMessage::Sentinel);
744                        }
745                    }
746                }
747            }
748        };
749
750        if let Some(ctrl) = closing_msg
751            && let Ok(bytes) = serde_json::to_vec(&ctrl)
752            && let Err(err) = framed_writer
753                .send(TwoPartMessage::from_header(bytes.into()))
754                .await
755        {
756            tracing::trace!(?err, ?ctrl, "request-stream closing-frame send failed");
757        }
758
759        let mut inner = framed_writer.into_inner();
760        if let Err(err) = inner.flush().await {
761            tracing::trace!(?err, "request-stream socket flush failed");
762        }
763        if let Err(err) = inner.shutdown().await {
764            tracing::trace!(?err, "request-stream socket shutdown failed");
765        }
766    }
767
768    async fn process_response_stream(
769        subject: String,
770        state: Arc<Mutex<State>>,
771        mut reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
772        writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
773    ) -> Result<()> {
774        let response_stream = {
775            let mut guard = state.lock().await;
776            let conn = guard
777                .rx_subjects
778                .remove(&subject)
779                .ok_or(error!("Subject not found: {}; upstream publisher specified a subject unknown to the downsteam subscriber", subject))?;
780            if let Some(key) = guard.subject_instance.remove(&subject)
781                && let Some(subjects) = guard.instance_subjects.get_mut(&key)
782            {
783                subjects.remove(&(StreamType::Response, subject.clone()));
784                if subjects.is_empty() {
785                    guard.instance_subjects.remove(&key);
786                }
787            }
788            conn
789        };
790
791        // unwrap response_stream
792        let RequestedRecvConnection {
793            context,
794            connection,
795            send_buffer_count,
796        } = response_stream;
797
798        // the [`Prologue`]
799        // there must be a second control message it indicate the other segment's generate method was successful
800        let prologue = reader
801            .next()
802            .await
803            .ok_or(error!("Connection closed without a ControlMessge"))??;
804
805        // deserialize prologue
806        let prologue = match prologue.into_message_type() {
807            TwoPartMessageType::HeaderOnly(header) => {
808                let prologue: ResponseStreamPrologue = serde_json::from_slice(&header)
809                    .map_err(|e| error!("Failed to deserialize ControlMessage: {}", e))?;
810                prologue
811            }
812            _ => {
813                // Worker sent a non-HeaderOnly frame in the prologue slot
814                // (protocol violation, version skew, corruption). Notify the
815                // requester so the generate call chain fails cleanly, then
816                // return Err so the connection task ends without panicking.
817                let msg = "malformed prologue: expected HeaderOnly ControlMessage";
818                let _ = connection.send(Err(msg.to_string()));
819                return Err(error!(msg));
820            }
821        };
822
823        // await the control message of GTG or Error, if error, then connection.send(Err(String)), which should fail the
824        // generate call chain
825        //
826        // note: this second control message might be delayed, but the expensive part of setting up the connection
827        // is both complete and ready for data flow; awaiting here is not a performance hit or problem and it allows
828        // us to trace the initial setup time vs the time to prologue
829        if let Some(error) = &prologue.error {
830            let _ = connection.send(Err(error.clone()));
831            return Err(error!("Received error prologue: {}", error));
832        }
833
834        // Buffer size is driven by the registration options
835        // ([`StreamOptions::send_buffer_count`]) rather than hard-coded; the
836        // same applies to `process_request_stream`. See #10293.
837        let (response_tx, response_rx) = data_plane_channel(send_buffer_count);
838
839        if connection
840            .send(Ok(crate::pipeline::network::StreamReceiver {
841                rx: response_rx,
842            }))
843            .is_err()
844        {
845            return Err(error!(
846                "The requester of the stream has been dropped before the connection was established"
847            ));
848        }
849
850        let (control_tx, control_rx) = mpsc::channel::<ControlMessage>(1);
851
852        // sender task
853        // issues control messages to the sender and when finished shuts down the socket
854        // this should be the last task to finish and must
855        let send_task = tokio::spawn(network_send_handler(writer, control_rx));
856
857        // forward task
858        let recv_task = tokio::spawn(network_receive_handler(
859            reader,
860            response_tx,
861            control_tx,
862            context.clone(),
863        ));
864
865        // check the results of each of the tasks
866        let (monitor_result, forward_result) = tokio::join!(send_task, recv_task);
867
868        monitor_result?;
869        forward_result?;
870
871        Ok(())
872    }
873
874    async fn network_receive_handler(
875        mut framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
876        response_tx: mpsc::Sender<Bytes>,
877        control_tx: mpsc::Sender<ControlMessage>,
878        context: Arc<dyn AsyncEngineContext>,
879    ) {
880        // loop over reading the tcp stream and checking if the writer is closed
881        let mut can_stop = true;
882        loop {
883            tokio::select! {
884                biased;
885
886                _ = response_tx.closed() => {
887                    tracing::trace!("response channel closed before the client finished writing data");
888                    let _ = control_tx.send(ControlMessage::Kill).await;
889                    break;
890                }
891
892                _ = context.killed() => {
893                    tracing::trace!("context kill signal received; shutting down");
894                    let _ = control_tx.send(ControlMessage::Kill).await;
895                    break;
896                }
897
898                _ = context.stopped(), if can_stop => {
899                    tracing::trace!("context stop signal received; shutting down");
900                    can_stop = false;
901                    let _ = control_tx.send(ControlMessage::Stop).await;
902                }
903
904                msg = framed_reader.next() => {
905                    match msg {
906                        Some(Ok(msg)) => {
907                            let (header, data) = msg.into_parts();
908
909                            // received a control message
910                            if !header.is_empty() {
911                                match process_control_message(header) {
912                                    Ok(ControlAction::Continue) => {}
913                                    Ok(ControlAction::Shutdown) => {
914                                        if !data.is_empty() {
915                                            // Sentinel-with-data is a protocol
916                                            // violation; kill this stream, don't
917                                            // assert!() the process down.
918                                            tracing::warn!(
919                                                data_len = data.len(),
920                                                "client sent Sentinel with data (protocol violation); killing stream"
921                                            );
922                                            let _ = control_tx.send(ControlMessage::Kill).await;
923                                            break;
924                                        }
925                                        tracing::trace!("received sentinel message; shutting down");
926                                        break;
927                                    }
928                                    Err(e) => {
929                                        // Malformed control message — kill only
930                                        // this stream.
931                                        tracing::warn!(err = ?e, "malformed control message, closing connection");
932                                        let _ = control_tx.send(ControlMessage::Kill).await;
933                                        break;
934                                    }
935                                }
936                            }
937
938                            if !data.is_empty()
939                                && let Err(err) = response_tx.send(data).await {
940                                    tracing::debug!(?err, "forwarding body/data to response channel failed");
941                                    let _ = control_tx.send(ControlMessage::Kill).await;
942                                    break;
943                                };
944                        }
945                        Some(Err(e)) => {
946                            // TCP RST or decode error from worker — kill only
947                            // this stream.
948                            tracing::warn!(err = ?e, "tcp stream read error from worker, closing connection");
949                            let _ = control_tx.send(ControlMessage::Kill).await;
950                            break;
951                        }
952                        None => {
953                            // this is allowed but we try to avoid it
954                            // the logic is that the client will tell us when its is done and the server
955                            // will close the connection naturally when the sentinel message is received
956                            // the client closing early represents a transport error outside the control of the
957                            // transport library
958                            tracing::trace!("tcp stream was closed by client");
959                            break;
960                        }
961                    }
962                }
963
964            }
965        }
966    }
967
968    async fn network_send_handler(
969        socket_tx: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
970        control_rx: mpsc::Receiver<ControlMessage>,
971    ) {
972        let mut socket_tx = socket_tx;
973        let mut control_rx = control_rx;
974
975        while let Some(control_msg) = control_rx.recv().await {
976            // Sentinel is a worker→frontend message; receiving one here means
977            // a producer is buggy. Skip rather than asserting — a stream-level
978            // bug must not panic the worker.
979            if matches!(control_msg, ControlMessage::Sentinel) {
980                tracing::warn!("received sentinel on send-side control channel; dropping");
981                continue;
982            }
983            let bytes = match serde_json::to_vec(&control_msg) {
984                Ok(b) => b,
985                Err(e) => {
986                    // Closed enum of small variants; serialization shouldn't
987                    // fail. If it ever does, log and skip rather than panic.
988                    tracing::warn!(err = ?e, ?control_msg, "failed to serialize control message");
989                    continue;
990                }
991            };
992            let message = TwoPartMessage::from_header(bytes.into());
993            match socket_tx.send(message).await {
994                Ok(_) => tracing::debug!(?control_msg, "issued control message"),
995                Err(e) => {
996                    tracing::debug!(err = ?e, ?control_msg, "failed to send control message")
997                }
998            }
999        }
1000
1001        let mut inner = socket_tx.into_inner();
1002        if let Err(e) = inner.flush().await {
1003            tracing::debug!("failed to flush socket: {e}");
1004        }
1005        if let Err(e) = inner.shutdown().await {
1006            tracing::debug!("failed to shutdown socket: {e}");
1007        }
1008    }
1009}
1010
1011enum ControlAction {
1012    Continue,
1013    Shutdown,
1014}
1015
1016fn process_control_message(message: Bytes) -> Result<ControlAction> {
1017    match serde_json::from_slice::<ControlMessage>(&message)? {
1018        ControlMessage::Sentinel => {
1019            // the client issued a sentinel message
1020            // it has finished writing data and is now awaiting the server to close the connection
1021            tracing::trace!("sentinel received; shutting down");
1022            Ok(ControlAction::Shutdown)
1023        }
1024        ControlMessage::Kill | ControlMessage::Stop => {
1025            // Worker→frontend control direction only carries Sentinel. Kill/Stop
1026            // here is a protocol violation; the caller turns this Err into a
1027            // stream-local Kill rather than a process-fatal event.
1028            anyhow::bail!("unexpected control message on response stream");
1029        }
1030    }
1031}
1032
1033#[cfg(test)]
1034mod tests {
1035    use super::*;
1036    use crate::engine::AsyncEngineContextProvider;
1037    use crate::pipeline::Context;
1038    use crate::pipeline::network::DEFAULT_SEND_BUFFER_COUNT;
1039    use tokio::io::{AsyncWriteExt, ReadHalf, WriteHalf};
1040    use tokio::net::TcpStream;
1041
1042    // Mock resolver that always fails to simulate the fallback scenario
1043    struct FailingIpResolver;
1044
1045    impl IpResolver for FailingIpResolver {
1046        fn local_ip(&self) -> Result<std::net::IpAddr, Error> {
1047            Err(Error::LocalIpAddressNotFound)
1048        }
1049
1050        fn local_ipv6(&self) -> Result<std::net::IpAddr, Error> {
1051            Err(Error::LocalIpAddressNotFound)
1052        }
1053    }
1054
1055    #[tokio::test]
1056    async fn test_tcp_stream_server_default_behavior() {
1057        // Test that TcpStreamServer::new works with default options
1058        // This verifies normal operation when IP detection succeeds
1059        let options = ServerOptions::default();
1060        let result = TcpStreamServer::new(options).await;
1061
1062        assert!(
1063            result.is_ok(),
1064            "TcpStreamServer::new should succeed with default options"
1065        );
1066
1067        let server = result.unwrap();
1068
1069        // Verify the server can be used by registering a stream
1070        let context = Context::new(());
1071        let stream_options = StreamOptions::builder()
1072            .context(context.context())
1073            .enable_request_stream(false)
1074            .enable_response_stream(true)
1075            .build()
1076            .unwrap();
1077
1078        let pending_connection = server.register(stream_options).await;
1079
1080        // Verify connection info is available and valid
1081        let connection_info = pending_connection
1082            .recv_stream
1083            .as_ref()
1084            .unwrap()
1085            .connection_info
1086            .clone();
1087
1088        let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1089        let socket_addr = tcp_info.address.parse::<std::net::SocketAddr>().unwrap();
1090
1091        // Should have a valid port assigned
1092        assert!(
1093            socket_addr.port() > 0,
1094            "Server should be assigned a valid port number"
1095        );
1096
1097        println!(
1098            "Server created successfully with address: {}",
1099            tcp_info.address
1100        );
1101    }
1102
1103    /// The data-plane channel helper sizes the mpsc buffer from
1104    /// `send_buffer_count` — this is the value `process_request_stream` /
1105    /// `process_response_stream` feed it. `max_capacity()` reflects the
1106    /// channel's configured buffer, so a custom value and the default both
1107    /// reach the channel. Guards against regressing back to a hard-coded 64.
1108    #[test]
1109    fn data_plane_channel_capacity_matches_send_buffer_count() {
1110        let (tx, _rx) = data_plane_channel::<()>(7);
1111        assert_eq!(tx.max_capacity(), 7);
1112
1113        let (tx, _rx) = data_plane_channel::<()>(DEFAULT_SEND_BUFFER_COUNT);
1114        assert_eq!(tx.max_capacity(), 64);
1115
1116        // A misconfigured 0 must clamp to 1, not panic (mpsc::channel(0) panics).
1117        let (tx, _rx) = data_plane_channel::<()>(0);
1118        assert_eq!(tx.max_capacity(), 1);
1119    }
1120
1121    /// `register` must thread `StreamOptions::send_buffer_count` through to the
1122    /// stored `RequestedSendConnection` / `RequestedRecvConnection` (the
1123    /// registration structs `process_*_stream` later destructure to size the
1124    /// channel). Verified here against the real registration path.
1125    #[tokio::test]
1126    async fn register_threads_send_buffer_count_into_connection_structs() {
1127        let server = TcpStreamServer::new(ServerOptions::default())
1128            .await
1129            .expect("server");
1130        let context = Context::new(());
1131        let options = StreamOptions::builder()
1132            .context(context.context())
1133            .enable_request_stream(true)
1134            .enable_response_stream(true)
1135            .send_buffer_count(7)
1136            .build()
1137            .unwrap();
1138
1139        let _pending = server.register(options).await;
1140
1141        let state = server.state.lock().await;
1142        assert_eq!(state.tx_subjects.len(), 1, "one request stream registered");
1143        assert_eq!(state.rx_subjects.len(), 1, "one response stream registered");
1144        assert!(
1145            state.tx_subjects.values().all(|c| c.send_buffer_count == 7),
1146            "send_buffer_count must reach RequestedSendConnection"
1147        );
1148        assert!(
1149            state.rx_subjects.values().all(|c| c.send_buffer_count == 7),
1150            "send_buffer_count must reach RequestedRecvConnection"
1151        );
1152    }
1153
1154    #[tokio::test]
1155    async fn test_tcp_stream_server_fallback_to_loopback() {
1156        // Test fallback behavior using a mock resolver that always fails
1157        // This guarantees the fallback logic is triggered
1158
1159        let options = ServerOptions::builder().port(0).build().unwrap();
1160
1161        // Use the failing resolver to force the fallback
1162        let result = TcpStreamServer::new_with_resolver(options, FailingIpResolver).await;
1163        assert!(
1164            result.is_ok(),
1165            "Server creation should succeed with fallback even when IP detection fails"
1166        );
1167
1168        let server = result.unwrap();
1169
1170        // Get the actual bound address by registering a stream
1171        let context = Context::new(());
1172        let stream_options = StreamOptions::builder()
1173            .context(context.context())
1174            .enable_request_stream(false)
1175            .enable_response_stream(true)
1176            .build()
1177            .unwrap();
1178
1179        let pending_connection = server.register(stream_options).await;
1180        let connection_info = pending_connection
1181            .recv_stream
1182            .as_ref()
1183            .unwrap()
1184            .connection_info
1185            .clone();
1186
1187        let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1188        let socket_addr = tcp_info.address.parse::<std::net::SocketAddr>().unwrap();
1189
1190        // With the failing resolver, fallback should ALWAYS be used
1191        let ip = socket_addr.ip();
1192        assert!(
1193            ip.is_loopback(),
1194            "Should use loopback when IP detection fails"
1195        );
1196
1197        // Verify it's specifically 127.0.0.1 (the fallback value from the patch)
1198        assert_eq!(
1199            ip,
1200            std::net::IpAddr::V4(std::net::Ipv4Addr::new(127, 0, 0, 1)),
1201            "Fallback should use exactly 127.0.0.1, got: {}",
1202            ip
1203        );
1204
1205        println!("SUCCESS: Fallback to 127.0.0.1 was confirmed: {}", ip);
1206
1207        // The server should work with the fallback IP
1208        assert!(socket_addr.port() > 0, "Server should have a valid port");
1209    }
1210
1211    /// Create a test server using the failing IP resolver (falls back to loopback).
1212    async fn test_server() -> Arc<TcpStreamServer> {
1213        TcpStreamServer::new_with_resolver(
1214            ServerOptions::builder().port(0).build().unwrap(),
1215            FailingIpResolver,
1216        )
1217        .await
1218        .unwrap()
1219    }
1220
1221    /// Helper: register a response stream and extract its subject string.
1222    async fn register_and_get_subject(
1223        server: &TcpStreamServer,
1224    ) -> (
1225        String,
1226        tokio::sync::oneshot::Receiver<Result<super::StreamReceiver, String>>,
1227    ) {
1228        let context = Context::new(());
1229        let options = StreamOptions::builder()
1230            .context(context.context())
1231            .enable_request_stream(false)
1232            .enable_response_stream(true)
1233            .build()
1234            .unwrap();
1235
1236        let pending = server.register(options).await;
1237        let recv_stream = pending.recv_stream.unwrap();
1238        let (conn_info, provider) = recv_stream.into_parts();
1239        let tcp_info: TcpStreamConnectionInfo = conn_info.try_into().unwrap();
1240        (tcp_info.subject, provider)
1241    }
1242
1243    /// Convenience constructor so tests don't repeat the struct literal.
1244    fn make_eid(
1245        namespace: &str,
1246        component: &str,
1247        endpoint: &str,
1248        instance_id: u64,
1249    ) -> EndpointInstanceId {
1250        EndpointInstanceId {
1251            namespace: namespace.to_string(),
1252            component: component.to_string(),
1253            endpoint: endpoint.to_string(),
1254            instance_id,
1255        }
1256    }
1257
1258    /// Helper: register a bidirectional pair (both request + response halves)
1259    /// and return both subjects + their providers.
1260    async fn register_and_get_bidi_subjects(
1261        server: &TcpStreamServer,
1262    ) -> (
1263        String,
1264        tokio::sync::oneshot::Receiver<Result<super::StreamSender, String>>,
1265        String,
1266        tokio::sync::oneshot::Receiver<Result<super::StreamReceiver, String>>,
1267    ) {
1268        let context = Context::new(());
1269        let options = StreamOptions::builder()
1270            .context(context.context())
1271            .enable_request_stream(true)
1272            .enable_response_stream(true)
1273            .build()
1274            .unwrap();
1275
1276        let pending = server.register(options).await;
1277        let send_stream = pending.send_stream.unwrap();
1278        let recv_stream = pending.recv_stream.unwrap();
1279        let (send_info, send_provider) = send_stream.into_parts();
1280        let (recv_info, recv_provider) = recv_stream.into_parts();
1281        let send_tcp_info: TcpStreamConnectionInfo = send_info.try_into().unwrap();
1282        let recv_tcp_info: TcpStreamConnectionInfo = recv_info.try_into().unwrap();
1283        (
1284            send_tcp_info.subject,
1285            send_provider,
1286            recv_tcp_info.subject,
1287            recv_provider,
1288        )
1289    }
1290
1291    /// `cancel_instance_streams` must drop the request-stream oneshot too,
1292    /// not just the response-stream one. Without the tagged tracker this test
1293    /// would hang on `send_provider.await` because the tx_subjects entry
1294    /// would leak past instance removal.
1295    #[tokio::test]
1296    async fn test_cancel_instance_streams_drops_both_bidi_halves() {
1297        let server = test_server().await;
1298        let (send_subj, send_provider, recv_subj, recv_provider) =
1299            register_and_get_bidi_subjects(&server).await;
1300
1301        let id = make_eid("ns", "comp", "generate", 7);
1302        assert!(
1303            server
1304                .associate_instance(&recv_subj, Some(&send_subj), &id)
1305                .await,
1306            "fresh instance must not be tombstoned"
1307        );
1308
1309        let cancelled = server.cancel_instance_streams(&id).await;
1310        assert_eq!(cancelled, 2, "both request + response halves must count");
1311
1312        assert!(
1313            recv_provider.await.is_err(),
1314            "recv provider should resolve with RecvError"
1315        );
1316        assert!(
1317            send_provider.await.is_err(),
1318            "send provider should resolve with RecvError after instance cancellation"
1319        );
1320    }
1321
1322    /// Pre-tombstoning an instance must drop both halves of a later
1323    /// `associate_instance(recv, Some(send), id)` call, not just the recv.
1324    #[tokio::test]
1325    async fn test_associate_instance_tombstone_cancels_both_bidi_halves() {
1326        let server = test_server().await;
1327        let id = make_eid("ns", "comp", "generate", 8);
1328        // Pre-tombstone the instance.
1329        server.cancel_instance_streams(&id).await;
1330
1331        let (send_subj, send_provider, recv_subj, recv_provider) =
1332            register_and_get_bidi_subjects(&server).await;
1333
1334        assert!(
1335            !server
1336                .associate_instance(&recv_subj, Some(&send_subj), &id)
1337                .await,
1338            "tombstoned instance must reject association"
1339        );
1340
1341        assert!(recv_provider.await.is_err());
1342        assert!(send_provider.await.is_err());
1343    }
1344
1345    #[tokio::test]
1346    async fn test_cancel_instance_streams_unblocks_receiver() {
1347        let server = test_server().await;
1348
1349        let (subject, provider) = register_and_get_subject(&server).await;
1350
1351        let id = make_eid("ns", "comp", "generate", 42);
1352        assert!(server.associate_instance(&subject, None, &id).await);
1353
1354        let cancelled = server.cancel_instance_streams(&id).await;
1355        assert_eq!(cancelled, 1);
1356
1357        // The oneshot receiver should now resolve with an error (sender dropped)
1358        let result = provider.await;
1359        assert!(result.is_err(), "Expected RecvError after cancellation");
1360    }
1361
1362    #[tokio::test]
1363    async fn test_cancel_instance_streams_multiple_subjects() {
1364        let server = test_server().await;
1365
1366        let (subj1, prov1) = register_and_get_subject(&server).await;
1367        let (subj2, prov2) = register_and_get_subject(&server).await;
1368        let (subj3, prov3) = register_and_get_subject(&server).await;
1369
1370        let id10 = make_eid("ns", "comp", "generate", 10);
1371        let id20 = make_eid("ns", "comp", "generate", 20);
1372
1373        // Associate first two with instance 10, third with instance 20
1374        assert!(server.associate_instance(&subj1, None, &id10).await);
1375        assert!(server.associate_instance(&subj2, None, &id10).await);
1376        assert!(server.associate_instance(&subj3, None, &id20).await);
1377
1378        // Cancel instance 10 -- should cancel 2 subjects
1379        let cancelled = server.cancel_instance_streams(&id10).await;
1380        assert_eq!(cancelled, 2);
1381
1382        assert!(prov1.await.is_err());
1383        assert!(prov2.await.is_err());
1384
1385        // Instance 20 should be unaffected -- cancel it separately
1386        let cancelled = server.cancel_instance_streams(&id20).await;
1387        assert_eq!(cancelled, 1);
1388        assert!(prov3.await.is_err());
1389    }
1390
1391    #[tokio::test]
1392    async fn test_cancel_instance_streams_nonexistent_instance() {
1393        let server = test_server().await;
1394
1395        let id = make_eid("ns", "comp", "generate", 999);
1396        let cancelled = server.cancel_instance_streams(&id).await;
1397        assert_eq!(cancelled, 0);
1398    }
1399
1400    #[tokio::test]
1401    async fn test_cancel_recv_stream_cleans_up_instance_tracking() {
1402        let server = test_server().await;
1403
1404        let (subject, _provider) = register_and_get_subject(&server).await;
1405        let id = make_eid("ns", "comp", "generate", 42);
1406        assert!(server.associate_instance(&subject, None, &id).await);
1407
1408        // Cancel the individual subject
1409        server.cancel_recv_stream(&subject).await;
1410
1411        // Instance should have no remaining subjects
1412        let cancelled = server.cancel_instance_streams(&id).await;
1413        assert_eq!(
1414            cancelled, 0,
1415            "Instance tracking should have been cleaned up"
1416        );
1417    }
1418
1419    #[tokio::test]
1420    async fn test_registered_stream_drop_runs_cleanup() {
1421        let server = test_server().await;
1422
1423        // Register a response stream but DON'T call into_parts -- just drop it
1424        let context = Context::new(());
1425        let options = StreamOptions::builder()
1426            .context(context.context())
1427            .enable_request_stream(false)
1428            .enable_response_stream(true)
1429            .build()
1430            .unwrap();
1431
1432        let pending = server.register(options).await;
1433        let recv_stream = pending.recv_stream.unwrap();
1434
1435        // Get the subject before dropping
1436        let tcp_info: TcpStreamConnectionInfo =
1437            recv_stream.connection_info.clone().try_into().unwrap();
1438        let subject = tcp_info.subject.clone();
1439
1440        // Verify it's in rx_subjects
1441        {
1442            let state = server.state.lock().await;
1443            assert!(state.rx_subjects.contains_key(&subject));
1444        }
1445
1446        // Drop the RegisteredStream -- RAII cleanup should fire
1447        drop(recv_stream);
1448
1449        // Give the spawned cleanup task a moment to run
1450        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1451
1452        // Verify it's been removed from rx_subjects
1453        {
1454            let state = server.state.lock().await;
1455            assert!(
1456                !state.rx_subjects.contains_key(&subject),
1457                "RAII cleanup should have removed the rx_subjects entry"
1458            );
1459        }
1460    }
1461
1462    #[tokio::test]
1463    async fn test_registered_stream_into_parts_disarms_cleanup() {
1464        let server = test_server().await;
1465
1466        let context = Context::new(());
1467        let options = StreamOptions::builder()
1468            .context(context.context())
1469            .enable_request_stream(false)
1470            .enable_response_stream(true)
1471            .build()
1472            .unwrap();
1473
1474        let pending = server.register(options).await;
1475        let recv_stream = pending.recv_stream.unwrap();
1476
1477        let tcp_info: TcpStreamConnectionInfo =
1478            recv_stream.connection_info.clone().try_into().unwrap();
1479        let subject = tcp_info.subject.clone();
1480
1481        // Call into_parts to disarm the cleanup
1482        let (_conn_info, _provider) = recv_stream.into_parts();
1483
1484        // Give any potential cleanup a moment to run
1485        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1486
1487        // The entry should still be in rx_subjects (cleanup was disarmed)
1488        {
1489            let state = server.state.lock().await;
1490            assert!(
1491                state.rx_subjects.contains_key(&subject),
1492                "into_parts() should disarm the RAII cleanup"
1493            );
1494        }
1495    }
1496
1497    #[tokio::test]
1498    async fn test_associate_after_cancel_is_immediately_cancelled() {
1499        // Simulates the race: cancel_instance_streams fires before associate_instance.
1500        let server = test_server().await;
1501
1502        let id = make_eid("ns", "comp", "generate", 42);
1503
1504        // Cancel BEFORE any subject is registered (tombstone).
1505        let cancelled = server.cancel_instance_streams(&id).await;
1506        assert_eq!(cancelled, 0);
1507
1508        // Now register a subject and try to associate it with the tombstoned instance.
1509        let (subject, provider) = register_and_get_subject(&server).await;
1510        let associated = server.associate_instance(&subject, None, &id).await;
1511
1512        // associate_instance should return false when the instance is tombstoned.
1513        assert!(
1514            !associated,
1515            "associate_instance on a tombstoned instance should return false"
1516        );
1517
1518        // The provider should resolve with an error because associate_instance
1519        // found the tombstone and immediately cancelled the subject.
1520        let result = provider.await;
1521        assert!(
1522            result.is_err(),
1523            "Late associate_instance on a tombstoned instance should immediately cancel"
1524        );
1525    }
1526
1527    #[tokio::test]
1528    async fn test_clear_tombstone_allows_new_associations() {
1529        let server = test_server().await;
1530
1531        let id = make_eid("ns", "comp", "generate", 42);
1532
1533        server.cancel_instance_streams(&id).await;
1534        server.clear_instance_tombstone(&id).await;
1535
1536        // Now associate should work normally (subject NOT cancelled).
1537        let (subject, _provider) = register_and_get_subject(&server).await;
1538        assert!(server.associate_instance(&subject, None, &id).await);
1539
1540        // Subject should be tracked, not cancelled.
1541        let cancelled = server.cancel_instance_streams(&id).await;
1542        assert_eq!(
1543            cancelled, 1,
1544            "After clearing tombstone, subjects should be tracked normally"
1545        );
1546    }
1547
1548    #[tokio::test]
1549    async fn test_cancel_does_not_affect_sibling_endpoint() {
1550        // Regression: cancelling "generate" must not cancel "prefill" subjects
1551        // that share the same instance_id (same backend runtime).
1552        let server = test_server().await;
1553
1554        let (gen_subj, gen_prov) = register_and_get_subject(&server).await;
1555        let (pre_subj, pre_prov) = register_and_get_subject(&server).await;
1556
1557        let gen_id = make_eid("ns", "comp", "generate", 42);
1558        let pre_id = make_eid("ns", "comp", "prefill", 42);
1559
1560        assert!(server.associate_instance(&gen_subj, None, &gen_id).await);
1561        assert!(server.associate_instance(&pre_subj, None, &pre_id).await);
1562
1563        // Cancel only the "generate" endpoint's subjects.
1564        let cancelled = server.cancel_instance_streams(&gen_id).await;
1565        assert_eq!(
1566            cancelled, 1,
1567            "Only the generate subject should be cancelled"
1568        );
1569        assert!(gen_prov.await.is_err());
1570
1571        // prefill must still be tracked.
1572        let still_pending = server.cancel_instance_streams(&pre_id).await;
1573        assert_eq!(still_pending, 1, "prefill subject should still be tracked");
1574        assert!(pre_prov.await.is_err());
1575    }
1576
1577    #[tokio::test]
1578    async fn test_tombstone_is_endpoint_scoped() {
1579        // Tombstoning "generate" must not prevent new associations on "prefill"
1580        // for the same instance_id.
1581        let server = test_server().await;
1582
1583        let gen_id = make_eid("ns", "comp", "generate", 42);
1584        let pre_id = make_eid("ns", "comp", "prefill", 42);
1585
1586        server.cancel_instance_streams(&gen_id).await;
1587
1588        // A new subject for "generate" should be rejected.
1589        let (gen_subj, gen_prov) = register_and_get_subject(&server).await;
1590        assert!(
1591            !server.associate_instance(&gen_subj, None, &gen_id).await,
1592            "generate should be tombstoned"
1593        );
1594        assert!(gen_prov.await.is_err());
1595
1596        // A new subject for "prefill" with the same instance_id should be accepted.
1597        let (pre_subj, _pre_prov) = register_and_get_subject(&server).await;
1598        assert!(
1599            server.associate_instance(&pre_subj, None, &pre_id).await,
1600            "prefill tombstone is independent; subject should be tracked"
1601        );
1602        let count = server.cancel_instance_streams(&pre_id).await;
1603        assert_eq!(count, 1, "prefill subject should be tracked normally");
1604    }
1605
1606    #[tokio::test]
1607    async fn test_cancel_does_not_affect_different_component() {
1608        // Regression: two services with different (namespace, component) but the
1609        // same endpoint name and the same pod-backed instance_id must not interfere,
1610        // even though they share a single TcpStreamServer runtime.
1611        let server = test_server().await;
1612
1613        let (subj_a, prov_a) = register_and_get_subject(&server).await;
1614        let (subj_b, prov_b) = register_and_get_subject(&server).await;
1615
1616        // Same endpoint name + instance_id, different namespace/component.
1617        let id_a = make_eid("ns-a", "comp-a", "generate", 42);
1618        let id_b = make_eid("ns-b", "comp-b", "generate", 42);
1619
1620        assert!(server.associate_instance(&subj_a, None, &id_a).await);
1621        assert!(server.associate_instance(&subj_b, None, &id_b).await);
1622
1623        // Cancel service A -- only subj_a should be affected.
1624        let cancelled = server.cancel_instance_streams(&id_a).await;
1625        assert_eq!(cancelled, 1, "Only service-A subject should be cancelled");
1626        assert!(prov_a.await.is_err());
1627
1628        // Service B subject must still be pending.
1629        let still_tracked = server.cancel_instance_streams(&id_b).await;
1630        assert_eq!(still_tracked, 1, "Service-B subject should be unaffected");
1631        assert!(prov_b.await.is_err());
1632    }
1633
1634    #[tokio::test(start_paused = true)]
1635    async fn test_tombstone_expires_after_ttl() {
1636        // After TOMBSTONE_TTL elapses, a previously-tombstoned identity must
1637        // accept new associations again, AND the entry must be physically
1638        // pruned from `removed_instances` so the set remains bounded.
1639        let server = test_server().await;
1640
1641        let id = make_eid("ns", "comp", "generate", 42);
1642
1643        // Tombstone the identity.
1644        server.cancel_instance_streams(&id).await;
1645        {
1646            let state = server.state.lock().await;
1647            assert!(state.removed_instances.contains_key(&id));
1648        }
1649
1650        // Advance past the TTL.
1651        tokio::time::advance(TOMBSTONE_TTL + Duration::from_secs(1)).await;
1652
1653        // associate_instance for the same identity should now succeed (no
1654        // longer tombstoned). Any new subject must be tracked normally.
1655        let (subject, _provider) = register_and_get_subject(&server).await;
1656        assert!(
1657            server.associate_instance(&subject, None, &id).await,
1658            "tombstone older than TTL should not block association"
1659        );
1660
1661        // The expired tombstone must have been pruned (lazy pruning fires on
1662        // every associate_instance/cancel_instance_streams call).
1663        {
1664            let state = server.state.lock().await;
1665            assert!(
1666                !state.removed_instances.contains_key(&id),
1667                "expired tombstone should be pruned, not retained"
1668            );
1669        }
1670    }
1671
1672    #[tokio::test(start_paused = true)]
1673    async fn test_tombstone_within_ttl_blocks_associate() {
1674        // Regression net for the original tombstone fix: a tombstone younger
1675        // than TTL must still cancel late-arriving associate_instance() calls.
1676        let server = test_server().await;
1677
1678        let id = make_eid("ns", "comp", "generate", 42);
1679        server.cancel_instance_streams(&id).await;
1680
1681        // Advance only a small fraction of the TTL.
1682        tokio::time::advance(Duration::from_secs(1)).await;
1683
1684        let (subject, provider) = register_and_get_subject(&server).await;
1685        assert!(
1686            !server.associate_instance(&subject, None, &id).await,
1687            "tombstone within TTL must still block association"
1688        );
1689        assert!(provider.await.is_err());
1690    }
1691
1692    #[tokio::test(start_paused = true)]
1693    async fn test_tombstone_lazy_prune_on_cancel() {
1694        // Old tombstones must be pruned on the next cancel_instance_streams
1695        // call, regardless of which identity is being tombstoned.
1696        let server = test_server().await;
1697
1698        let id_old = make_eid("ns", "comp", "generate", 1);
1699        let id_new = make_eid("ns", "comp", "generate", 2);
1700
1701        server.cancel_instance_streams(&id_old).await;
1702        tokio::time::advance(TOMBSTONE_TTL + Duration::from_secs(1)).await;
1703        server.cancel_instance_streams(&id_new).await;
1704
1705        let state = server.state.lock().await;
1706        assert!(
1707            !state.removed_instances.contains_key(&id_old),
1708            "old tombstone should be pruned by the next cancel_instance_streams call"
1709        );
1710        assert!(
1711            state.removed_instances.contains_key(&id_new),
1712            "fresh tombstone should be retained"
1713        );
1714        assert_eq!(state.removed_instances.len(), 1);
1715    }
1716
1717    #[tokio::test]
1718    async fn test_clear_tombstone_only_affects_named_identity() {
1719        // Documents the monotonic-lease invariant: `clear_instance_tombstone`
1720        // for one EndpointInstanceId must not touch a sibling entry. With etcd
1721        // lease IDs this defensive code rarely fires (new lease = new
1722        // EndpointInstanceId), but the per-key scope must hold.
1723        let server = test_server().await;
1724
1725        let id_a = make_eid("ns", "comp", "generate", 1);
1726        let id_b = make_eid("ns", "comp", "generate", 2);
1727
1728        server.cancel_instance_streams(&id_a).await;
1729        server.clear_instance_tombstone(&id_b).await;
1730
1731        let state = server.state.lock().await;
1732        assert!(
1733            state.removed_instances.contains_key(&id_a),
1734            "clearing a different identity must not remove id_a's tombstone"
1735        );
1736    }
1737
1738    #[tokio::test]
1739    async fn test_tombstone_scoped_to_full_identity() {
1740        // A tombstone on (ns-a, comp-a, generate, 42) must not block
1741        // associations on (ns-b, comp-b, generate, 42).
1742        let server = test_server().await;
1743
1744        let id_a = make_eid("ns-a", "comp-a", "generate", 42);
1745        let id_b = make_eid("ns-b", "comp-b", "generate", 42);
1746
1747        // Tombstone only service A.
1748        server.cancel_instance_streams(&id_a).await;
1749
1750        // Service A is tombstoned — new association is rejected.
1751        let (subj_a, prov_a) = register_and_get_subject(&server).await;
1752        assert!(!server.associate_instance(&subj_a, None, &id_a).await);
1753        assert!(prov_a.await.is_err());
1754
1755        // Service B with same endpoint name + instance_id must be accepted.
1756        let (subj_b, _prov_b) = register_and_get_subject(&server).await;
1757        assert!(
1758            server.associate_instance(&subj_b, None, &id_b).await,
1759            "Different namespace/component must not be tombstoned"
1760        );
1761        assert_eq!(server.cancel_instance_streams(&id_b).await, 1);
1762    }
1763
1764    type TestFramedRead = FramedRead<ReadHalf<TcpStream>, TwoPartCodec>;
1765    type TestFramedWrite = FramedWrite<WriteHalf<TcpStream>, TwoPartCodec>;
1766    type TestResponseStream = (TestFramedRead, TestFramedWrite, StreamReceiver);
1767
1768    /// Stand up a TcpStreamServer, register a response stream, connect a
1769    /// client, drive the handshake + prologue, and return the client-side
1770    /// framed reader/writer along with the receiver.
1771    async fn open_registered_response_stream() -> TestResponseStream {
1772        let options = ServerOptions::builder().port(0).build().unwrap();
1773        let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver)
1774            .await
1775            .unwrap();
1776        let context = Context::new(());
1777        let stream_options = StreamOptions::builder()
1778            .context(context.context())
1779            .enable_request_stream(false)
1780            .enable_response_stream(true)
1781            .build()
1782            .unwrap();
1783        let pending_connection = server.register(stream_options).await;
1784        let registered_stream = pending_connection.recv_stream.unwrap();
1785        let (connection_info, stream_provider) = registered_stream.into_parts();
1786        let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1787
1788        let stream = TcpStream::connect(&tcp_info.address).await.unwrap();
1789        let (read_half, write_half) = tokio::io::split(stream);
1790        let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1791        let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
1792
1793        let handshake = CallHomeHandshake {
1794            subject: tcp_info.subject,
1795            stream_type: StreamType::Response,
1796        };
1797        framed_writer
1798            .send(TwoPartMessage::from_header(
1799                serde_json::to_vec(&handshake).unwrap().into(),
1800            ))
1801            .await
1802            .unwrap();
1803        framed_writer
1804            .send(TwoPartMessage::from_header(
1805                serde_json::to_vec(&ResponseStreamPrologue { error: None })
1806                    .unwrap()
1807                    .into(),
1808            ))
1809            .await
1810            .unwrap();
1811
1812        // SAFETY (test-only): healthy localhost handshake always resolves all
1813        // three layers; a panic here means the harness is broken.
1814        let receiver = tokio::time::timeout(std::time::Duration::from_secs(1), stream_provider)
1815            .await
1816            .expect("server should establish response stream within timeout")
1817            .expect("stream provider should not be dropped")
1818            .expect("response stream should be accepted");
1819
1820        (framed_reader, framed_writer, receiver)
1821    }
1822
1823    async fn recv_control_message(framed_reader: &mut TestFramedRead) -> ControlMessage {
1824        // SAFETY (test-only): a misbehaving server in any of these layers is
1825        // exactly the harness failure we want surfaced as a test panic.
1826        let message = tokio::time::timeout(std::time::Duration::from_secs(1), framed_reader.next())
1827            .await
1828            .expect("server should send a control message within timeout")
1829            .expect("server should not close before sending control")
1830            .expect("control message should decode");
1831        let (header, data) = message.optional_parts();
1832        assert!(data.is_none(), "control message should not contain data");
1833        serde_json::from_slice(header.expect("control header missing").as_ref()).unwrap()
1834    }
1835
1836    /// Sending an unexpected control message (Stop or Kill from the data
1837    /// direction) is a protocol violation. The server's
1838    /// network_receive_handler must reply with ControlMessage::Kill on
1839    /// that stream alone, not panic.
1840    #[tokio::test]
1841    async fn test_tcp_stream_server_sends_kill_on_unexpected_control_message() {
1842        let (mut framed_reader, mut framed_writer, _receiver) =
1843            open_registered_response_stream().await;
1844
1845        framed_writer
1846            .send(TwoPartMessage::from_header(
1847                serde_json::to_vec(&ControlMessage::Stop).unwrap().into(),
1848            ))
1849            .await
1850            .unwrap();
1851
1852        assert_eq!(
1853            recv_control_message(&mut framed_reader).await,
1854            ControlMessage::Kill,
1855            "unexpected control message should kill only this stream"
1856        );
1857    }
1858
1859    /// A framing/decode error from the worker side is unrecoverable for
1860    /// this stream but must not panic the worker. Server should send Kill
1861    /// and tear down only this connection.
1862    #[tokio::test]
1863    async fn test_tcp_stream_server_sends_kill_on_read_error() {
1864        let (mut framed_reader, framed_writer, _receiver) = open_registered_response_stream().await;
1865
1866        let mut raw_writer = framed_writer.into_inner();
1867        raw_writer.write_all(&[0u8; 8]).await.unwrap();
1868        raw_writer.shutdown().await.unwrap();
1869
1870        assert_eq!(
1871            recv_control_message(&mut framed_reader).await,
1872            ControlMessage::Kill,
1873            "framing read error should kill only this stream"
1874        );
1875    }
1876
1877    /// Sentinel is supposed to be header-only. A misbehaving client that
1878    /// attaches a data payload must not panic the worker via assert!().
1879    #[tokio::test]
1880    async fn test_tcp_stream_server_sends_kill_on_sentinel_with_data() {
1881        let (mut framed_reader, mut framed_writer, _receiver) =
1882            open_registered_response_stream().await;
1883
1884        let header = serde_json::to_vec(&ControlMessage::Sentinel)
1885            .unwrap()
1886            .into();
1887        framed_writer
1888            .send(TwoPartMessage::from_parts(
1889                header,
1890                Bytes::from_static(b"unexpected payload"),
1891            ))
1892            .await
1893            .unwrap();
1894
1895        assert_eq!(
1896            recv_control_message(&mut framed_reader).await,
1897            ControlMessage::Kill,
1898            "Sentinel with data should kill only this stream"
1899        );
1900    }
1901
1902    /// The prologue must be a HeaderOnly frame. A non-HeaderOnly prologue
1903    /// (data-only or mixed) must surface as Err to the requester rather
1904    /// than panic the worker.
1905    #[tokio::test]
1906    async fn test_tcp_stream_server_returns_error_on_invalid_prologue() {
1907        let options = ServerOptions::builder().port(0).build().unwrap();
1908        let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver)
1909            .await
1910            .unwrap();
1911        let context = Context::new(());
1912        let stream_options = StreamOptions::builder()
1913            .context(context.context())
1914            .enable_request_stream(false)
1915            .enable_response_stream(true)
1916            .build()
1917            .unwrap();
1918        let pending_connection = server.register(stream_options).await;
1919        let registered_stream = pending_connection.recv_stream.unwrap();
1920        let (connection_info, stream_provider) = registered_stream.into_parts();
1921        let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1922
1923        let stream = TcpStream::connect(&tcp_info.address).await.unwrap();
1924        let (_read_half, write_half) = tokio::io::split(stream);
1925        let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
1926
1927        let handshake = CallHomeHandshake {
1928            subject: tcp_info.subject,
1929            stream_type: StreamType::Response,
1930        };
1931        framed_writer
1932            .send(TwoPartMessage::from_header(
1933                serde_json::to_vec(&handshake).unwrap().into(),
1934            ))
1935            .await
1936            .unwrap();
1937
1938        // Send a data-only frame in the prologue slot.
1939        framed_writer
1940            .send(TwoPartMessage::from_data(Bytes::from_static(
1941                b"not a prologue",
1942            )))
1943            .await
1944            .unwrap();
1945
1946        let outcome = tokio::time::timeout(std::time::Duration::from_secs(1), stream_provider)
1947            .await
1948            .expect("stream provider should resolve quickly")
1949            .expect("stream provider channel should not be dropped");
1950        // StreamReceiver doesn't impl Debug, so we can't use `.expect_err`.
1951        match outcome {
1952            Err(err) => assert!(
1953                err.contains("malformed prologue"),
1954                "expected malformed-prologue error, got: {err}"
1955            ),
1956            Ok(_) => panic!("invalid prologue should produce an error, but got Ok"),
1957        }
1958    }
1959
1960    // ==================== request_stream_send_handler integration tests ====================
1961    //
1962    // These exercise the closing-message contract of `request_stream_send_handler`
1963    // end-to-end: register a request stream, dial it as a raw client (so we can
1964    // inspect frames directly), then trigger each of the exit branches and
1965    // assert which ControlMessage arrives on the wire.
1966
1967    use futures::SinkExt;
1968
1969    /// Register a request stream and dial it with a raw client. Returns the
1970    /// framed reader on the raw client side, the StreamSender held by the
1971    /// upstream, and the upstream's engine context (so the test can drive
1972    /// kill / stop externally).
1973    async fn register_and_dial_request_stream(
1974        server: &TcpStreamServer,
1975    ) -> (
1976        FramedRead<tokio::io::ReadHalf<TcpStream>, TwoPartCodec>,
1977        super::StreamSender,
1978        Arc<dyn AsyncEngineContext>,
1979    ) {
1980        let upstream_ctx = Context::new(()).context();
1981        let options = StreamOptions::builder()
1982            .context(upstream_ctx.clone())
1983            .enable_request_stream(true)
1984            .enable_response_stream(false)
1985            .build()
1986            .unwrap();
1987
1988        let pending = server.register(options).await;
1989        let send_stream = pending.send_stream.unwrap();
1990        let (conn_info, send_provider) = send_stream.into_parts();
1991        let tcp_info: TcpStreamConnectionInfo = conn_info.try_into().unwrap();
1992
1993        let raw = TcpStream::connect(&tcp_info.address).await.unwrap();
1994        let (read_half, write_half) = tokio::io::split(raw);
1995        let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1996        let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
1997
1998        let handshake = super::CallHomeHandshake {
1999            subject: tcp_info.subject.clone(),
2000            stream_type: StreamType::Request,
2001        };
2002        let handshake_bytes = serde_json::to_vec(&handshake).unwrap();
2003        framed_writer
2004            .send(TwoPartMessage::from_header(handshake_bytes.into()))
2005            .await
2006            .unwrap();
2007        drop(framed_writer);
2008
2009        let sender = send_provider.await.unwrap().unwrap();
2010        (framed_reader, sender, upstream_ctx)
2011    }
2012
2013    /// Pull frames off the raw client reader until the first `ControlMessage`
2014    /// arrives, ignoring any DataOnly frames before it. Returns the variant.
2015    async fn next_control_message(
2016        reader: &mut FramedRead<tokio::io::ReadHalf<TcpStream>, TwoPartCodec>,
2017    ) -> ControlMessage {
2018        loop {
2019            let frame = reader
2020                .next()
2021                .await
2022                .expect("socket closed before control message arrived")
2023                .expect("decode error");
2024            if let Some(header) = frame.header() {
2025                return serde_json::from_slice::<ControlMessage>(header)
2026                    .expect("invalid control message bytes");
2027            }
2028            // DataOnly frame — skip and keep reading.
2029        }
2030    }
2031
2032    /// Dropping the upstream's StreamSender drains `request_rx` and the server
2033    /// emits [`ControlMessage::Sentinel`] as the closing frame.
2034    #[tokio::test]
2035    async fn test_request_stream_sends_sentinel_on_clean_drop() {
2036        let server = test_server().await;
2037        let (mut reader, sender, _ctx) = register_and_dial_request_stream(&server).await;
2038
2039        drop(sender);
2040
2041        let ctrl = next_control_message(&mut reader).await;
2042        assert!(
2043            matches!(ctrl, ControlMessage::Sentinel),
2044            "clean drain should emit Sentinel, got {ctrl:?}"
2045        );
2046    }
2047
2048    /// `context.kill()` makes the server emit [`ControlMessage::Kill`] before
2049    /// shutting down the write half.
2050    #[tokio::test]
2051    async fn test_request_stream_sends_kill_on_context_killed() {
2052        let server = test_server().await;
2053        let (mut reader, _sender, ctx) = register_and_dial_request_stream(&server).await;
2054
2055        ctx.kill();
2056
2057        let ctrl = next_control_message(&mut reader).await;
2058        assert!(
2059            matches!(ctrl, ControlMessage::Kill),
2060            "context.kill() should emit Kill, got {ctrl:?}"
2061        );
2062    }
2063
2064    /// `context.stop()` makes the server emit [`ControlMessage::Stop`] before
2065    /// shutting down the write half.
2066    #[tokio::test]
2067    async fn test_request_stream_sends_stop_on_context_stopped() {
2068        let server = test_server().await;
2069        let (mut reader, _sender, ctx) = register_and_dial_request_stream(&server).await;
2070
2071        ctx.stop();
2072
2073        let ctrl = next_control_message(&mut reader).await;
2074        assert!(
2075            matches!(ctrl, ControlMessage::Stop),
2076            "context.stop() should emit Stop, got {ctrl:?}"
2077        );
2078    }
2079}