1use 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
15const 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
50pub trait IpResolver {
52 fn local_ip(&self) -> Result<std::net::IpAddr, Error>;
53 fn local_ipv6(&self) -> Result<std::net::IpAddr, Error>;
54}
55
56pub 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
87pub struct TcpStreamServer {
91 local_ip: String,
92 local_port: u16,
93 state: Arc<Mutex<State>>,
94}
95
96#[allow(dead_code)]
103struct RequestedSendConnection {
104 context: Arc<dyn AsyncEngineContext>,
105 connection: oneshot::Sender<Result<StreamSender, String>>,
106 send_buffer_count: usize,
109}
110
111struct RequestedRecvConnection {
112 context: Arc<dyn AsyncEngineContext>,
113 connection: oneshot::Sender<Result<StreamReceiver, String>>,
114 send_buffer_count: usize,
117}
118
119fn data_plane_channel<T>(send_buffer_count: usize) -> (mpsc::Sender<T>, mpsc::Receiver<T>) {
125 mpsc::channel(send_buffer_count.max(1))
130}
131
132#[derive(Default)]
149struct State {
150 tx_subjects: HashMap<String, RequestedSendConnection>,
151 rx_subjects: HashMap<String, RequestedRecvConnection>,
152 subject_instance: HashMap<String, EndpointInstanceId>,
155 instance_subjects: HashMap<EndpointInstanceId, HashSet<(StreamType, String)>>,
160 removed_instances: HashMap<EndpointInstanceId, Instant>,
164 handle: Option<tokio::task::JoinHandle<Result<()>>>,
165}
166
167fn 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 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 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 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 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 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 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 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#[async_trait::async_trait]
383impl ResponseService for TcpStreamServer {
384 async fn register(&self, options: StreamOptions) -> PendingConnections {
405 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 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 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
514async 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 let (stream, _addr) = match listener.accept().await {
555 Ok((stream, _addr)) => (stream, _addr),
556 Err(e) => {
557 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 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 async fn process_stream(stream: tokio::net::TcpStream, state: Arc<Mutex<State>>) -> Result<()> {
598 let (read_half, write_half) = tokio::io::split(stream);
600
601 let mut framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
603 let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
604
605 let first_message = framed_reader
608 .next()
609 .await
610 .ok_or(error!("Connection closed without a ControlMessage"))??;
611
612 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 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 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 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 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 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 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 let RequestedRecvConnection {
793 context,
794 connection,
795 send_buffer_count,
796 } = response_stream;
797
798 let prologue = reader
801 .next()
802 .await
803 .ok_or(error!("Connection closed without a ControlMessge"))??;
804
805 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 let msg = "malformed prologue: expected HeaderOnly ControlMessage";
818 let _ = connection.send(Err(msg.to_string()));
819 return Err(error!(msg));
820 }
821 };
822
823 if let Some(error) = &prologue.error {
830 let _ = connection.send(Err(error.clone()));
831 return Err(error!("Received error prologue: {}", error));
832 }
833
834 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 let send_task = tokio::spawn(network_send_handler(writer, control_rx));
856
857 let recv_task = tokio::spawn(network_receive_handler(
859 reader,
860 response_tx,
861 control_tx,
862 context.clone(),
863 ));
864
865 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 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 if !header.is_empty() {
911 match process_control_message(header) {
912 Ok(ControlAction::Continue) => {}
913 Ok(ControlAction::Shutdown) => {
914 if !data.is_empty() {
915 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 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 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 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 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 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 tracing::trace!("sentinel received; shutting down");
1022 Ok(ControlAction::Shutdown)
1023 }
1024 ControlMessage::Kill | ControlMessage::Stop => {
1025 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 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 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 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 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 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 #[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 let (tx, _rx) = data_plane_channel::<()>(0);
1118 assert_eq!(tx.max_capacity(), 1);
1119 }
1120
1121 #[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 let options = ServerOptions::builder().port(0).build().unwrap();
1160
1161 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 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 let ip = socket_addr.ip();
1192 assert!(
1193 ip.is_loopback(),
1194 "Should use loopback when IP detection fails"
1195 );
1196
1197 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 assert!(socket_addr.port() > 0, "Server should have a valid port");
1209 }
1210
1211 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 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 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 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 #[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 #[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 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 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 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 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 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 server.cancel_recv_stream(&subject).await;
1410
1411 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 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 let tcp_info: TcpStreamConnectionInfo =
1437 recv_stream.connection_info.clone().try_into().unwrap();
1438 let subject = tcp_info.subject.clone();
1439
1440 {
1442 let state = server.state.lock().await;
1443 assert!(state.rx_subjects.contains_key(&subject));
1444 }
1445
1446 drop(recv_stream);
1448
1449 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1451
1452 {
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 let (_conn_info, _provider) = recv_stream.into_parts();
1483
1484 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1486
1487 {
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 let server = test_server().await;
1501
1502 let id = make_eid("ns", "comp", "generate", 42);
1503
1504 let cancelled = server.cancel_instance_streams(&id).await;
1506 assert_eq!(cancelled, 0);
1507
1508 let (subject, provider) = register_and_get_subject(&server).await;
1510 let associated = server.associate_instance(&subject, None, &id).await;
1511
1512 assert!(
1514 !associated,
1515 "associate_instance on a tombstoned instance should return false"
1516 );
1517
1518 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 let (subject, _provider) = register_and_get_subject(&server).await;
1538 assert!(server.associate_instance(&subject, None, &id).await);
1539
1540 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 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 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 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 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 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 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 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 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 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 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 let server = test_server().await;
1640
1641 let id = make_eid("ns", "comp", "generate", 42);
1642
1643 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 tokio::time::advance(TOMBSTONE_TTL + Duration::from_secs(1)).await;
1652
1653 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 {
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 let server = test_server().await;
1677
1678 let id = make_eid("ns", "comp", "generate", 42);
1679 server.cancel_instance_streams(&id).await;
1680
1681 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 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 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 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 server.cancel_instance_streams(&id_a).await;
1749
1750 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 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 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 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 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 #[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 #[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 #[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 #[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 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 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 use futures::SinkExt;
1968
1969 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 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 }
2030 }
2031
2032 #[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 #[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 #[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}