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::time::Instant;
13
14const TOMBSTONE_TTL: Duration = Duration::from_secs(5);
19
20use bytes::Bytes;
21use derive_builder::Builder;
22use futures::{SinkExt, StreamExt};
23use local_ip_address::{Error, list_afinet_netifas, local_ip, local_ipv6};
24use parking_lot::Mutex;
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();
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();
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();
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();
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();
361 state.removed_instances.remove(id);
362 }
363
364 async fn start(local_ip: String, local_port: u16, state: Arc<Mutex<State>>) -> Result<u16> {
365 let addr = format!("{}:{}", local_ip, local_port);
366 let state_clone = state.clone();
367 let (ready_tx, ready_rx) = tokio::sync::oneshot::channel::<Result<u16>>();
368 {
369 let mut guard = state.lock();
370 if guard.handle.is_some() {
371 panic!("TcpStreamServer already started");
372 }
373 guard.handle = Some(tokio::spawn(tcp_listener(addr, state_clone, ready_tx)));
374 }
375 let local_port = ready_rx.await??;
376 Ok(local_port)
377 }
378
379 fn insert_request_stream(&self, subject: String, connection: RequestedSendConnection) {
380 self.state.lock().tx_subjects.insert(subject, connection);
381 }
382
383 fn insert_response_stream(&self, subject: String, connection: RequestedRecvConnection) {
384 self.state.lock().rx_subjects.insert(subject, connection);
385 }
386
387 fn take_request_stream(state: &Mutex<State>, subject: &str) -> Option<RequestedSendConnection> {
388 let mut state = state.lock();
389 let connection = state.tx_subjects.remove(subject);
390 if let Some(key) = state.subject_instance.remove(subject)
391 && let Some(subjects) = state.instance_subjects.get_mut(&key)
392 {
393 subjects.remove(&(StreamType::Request, subject.to_string()));
394 if subjects.is_empty() {
395 state.instance_subjects.remove(&key);
396 }
397 }
398 connection
399 }
400
401 fn take_response_stream(
402 state: &Mutex<State>,
403 subject: &str,
404 ) -> Option<RequestedRecvConnection> {
405 let mut state = state.lock();
406 let connection = state.rx_subjects.remove(subject);
407 if let Some(key) = state.subject_instance.remove(subject)
408 && let Some(subjects) = state.instance_subjects.get_mut(&key)
409 {
410 subjects.remove(&(StreamType::Response, subject.to_string()));
411 if subjects.is_empty() {
412 state.instance_subjects.remove(&key);
413 }
414 }
415 connection
416 }
417}
418
419#[async_trait::async_trait]
421impl ResponseService for TcpStreamServer {
422 async fn register(&self, options: StreamOptions) -> PendingConnections {
443 let address = format!("{}:{}", self.local_ip, self.local_port);
446 tracing::debug!("Registering new TcpStream on {address}");
447
448 let send_stream = if options.enable_request_stream {
449 let sender_subject = uuid::Uuid::new_v4().to_string();
450 let registry_subject = sender_subject.clone();
451
452 let (pending_sender_tx, pending_sender_rx) = oneshot::channel();
453
454 let connection_info = RequestedSendConnection {
455 context: options.context.clone(),
456 connection: pending_sender_tx,
457 send_buffer_count: options.send_buffer_count,
458 };
459
460 let cleanup_subject = sender_subject.clone();
461 let cleanup_state = self.state.clone();
462 let registered_stream = RegisteredStream::new(
463 TcpStreamConnectionInfo {
464 address: address.clone(),
465 subject: sender_subject,
466 context: options.context.id().to_string(),
467 stream_type: StreamType::Request,
468 }
469 .into(),
470 pending_sender_rx,
471 )
472 .with_cleanup(move || {
473 tokio::spawn(async move {
475 let mut state = cleanup_state.lock();
476 state.tx_subjects.remove(&cleanup_subject);
477 if let Some(key) = state.subject_instance.remove(&cleanup_subject)
478 && let Some(subjects) = state.instance_subjects.get_mut(&key)
479 {
480 subjects.remove(&(StreamType::Request, cleanup_subject.clone()));
481 if subjects.is_empty() {
482 state.instance_subjects.remove(&key);
483 }
484 }
485 });
486 });
487
488 self.insert_request_stream(registry_subject, connection_info);
489
490 Some(registered_stream)
491 } else {
492 None
493 };
494
495 let recv_stream = if options.enable_response_stream {
496 let (pending_recver_tx, pending_recver_rx) = oneshot::channel();
497 let receiver_subject = uuid::Uuid::new_v4().to_string();
498 let registry_subject = receiver_subject.clone();
499
500 let connection_info = RequestedRecvConnection {
501 context: options.context.clone(),
502 connection: pending_recver_tx,
503 send_buffer_count: options.send_buffer_count,
504 };
505
506 let cleanup_subject = receiver_subject.clone();
507 let cleanup_state = self.state.clone();
508 let registered_stream = RegisteredStream::new(
509 TcpStreamConnectionInfo {
510 address: address.clone(),
511 subject: receiver_subject,
512 context: options.context.id().to_string(),
513 stream_type: StreamType::Response,
514 }
515 .into(),
516 pending_recver_rx,
517 )
518 .with_cleanup(move || {
519 tokio::spawn(async move {
521 let mut state = cleanup_state.lock();
522 state.rx_subjects.remove(&cleanup_subject);
523 if let Some(key) = state.subject_instance.remove(&cleanup_subject)
524 && let Some(subjects) = state.instance_subjects.get_mut(&key)
525 {
526 subjects.remove(&(StreamType::Response, cleanup_subject.clone()));
527 if subjects.is_empty() {
528 state.instance_subjects.remove(&key);
529 }
530 }
531 });
532 });
533
534 self.insert_response_stream(registry_subject, connection_info);
535
536 Some(registered_stream)
537 } else {
538 None
539 };
540
541 PendingConnections {
542 send_stream,
543 recv_stream,
544 }
545 }
546}
547
548async fn tcp_listener(
555 addr: String,
556 state: Arc<Mutex<State>>,
557 read_tx: tokio::sync::oneshot::Sender<Result<u16>>,
558) -> Result<()> {
559 let listener = tokio::net::TcpListener::bind(&addr)
560 .await
561 .map_err(|e| anyhow::anyhow!("Failed to start TcpListender on {}: {}", addr, e));
562
563 let listener = match listener {
564 Ok(listener) => {
565 let addr = listener
566 .local_addr()
567 .map_err(|e| anyhow::anyhow!("Failed get SocketAddr: {:?}", e))
568 .unwrap();
569
570 read_tx
571 .send(Ok(addr.port()))
572 .expect("Failed to send ready signal");
573
574 listener
575 }
576 Err(e) => {
577 read_tx.send(Err(e)).expect("Failed to send ready signal");
578 return Err(anyhow::anyhow!("Failed to start TcpListender on {}", addr));
579 }
580 };
581
582 loop {
583 let (stream, _addr) = match listener.accept().await {
589 Ok((stream, _addr)) => (stream, _addr),
590 Err(e) => {
591 tracing::warn!("failed to accept tcp connection: {e}");
593 eprintln!("failed to accept tcp connection: {}", e);
594 continue;
595 }
596 };
597
598 match stream.set_nodelay(true) {
599 Ok(_) => (),
600 Err(e) => {
601 tracing::warn!("failed to set tcp stream to nodelay: {e}");
602 }
603 }
604
605 match stream.set_linger(Some(std::time::Duration::from_secs(0))) {
606 Ok(_) => (),
607 Err(e) => {
608 tracing::warn!("failed to set tcp stream to linger: {e}");
609 }
610 }
611
612 tokio::spawn(handle_connection(stream, state.clone()));
613 }
614
615 async fn handle_connection(stream: tokio::net::TcpStream, state: Arc<Mutex<State>>) {
618 let result = process_stream(stream, state).await;
619 match result {
620 Ok(_) => tracing::trace!("successfully processed tcp connection"),
621 Err(e) => {
622 tracing::warn!("failed to handle tcp connection: {e}");
623 #[cfg(debug_assertions)]
624 eprintln!("failed to handle tcp connection: {}", e);
625 }
626 }
627 }
628
629 async fn process_stream(stream: tokio::net::TcpStream, state: Arc<Mutex<State>>) -> Result<()> {
632 let (read_half, write_half) = tokio::io::split(stream);
634
635 let mut framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
637 let framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
638
639 let first_message = framed_reader
642 .next()
643 .await
644 .ok_or(error!("Connection closed without a ControlMessage"))??;
645
646 let handshake: CallHomeHandshake = match first_message.header() {
649 Some(header) => serde_json::from_slice(header).map_err(|e| {
650 error!(
651 "Failed to deserialize the first message as a valid `CallHomeHandshake`: {e}",
652 )
653 })?,
654 None => {
655 return Err(error!("Expected ControlMessage, got DataMessage"));
656 }
657 };
658
659 match handshake.stream_type {
661 StreamType::Request => {
662 process_request_stream(handshake.subject, state, framed_reader, framed_writer).await
663 }
664 StreamType::Response => {
665 process_response_stream(handshake.subject, state, framed_reader, framed_writer)
666 .await
667 }
668 }
669 }
670
671 async fn process_request_stream(
682 subject: String,
683 state: Arc<Mutex<State>>,
684 reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
685 writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
686 ) -> Result<()> {
687 drop(reader);
689
690 let request_stream = TcpStreamServer::take_request_stream(&state, &subject).ok_or_else(|| {
691 error!(
692 "Subject not found: {}; downstream subscriber specified a subject unknown to the upstream publisher",
693 subject
694 )
695 })?;
696
697 let RequestedSendConnection {
698 context,
699 connection,
700 send_buffer_count,
701 } = request_stream;
702
703 let (request_tx, request_rx) = data_plane_channel(send_buffer_count);
707
708 if connection
709 .send(Ok(crate::pipeline::network::StreamSender {
710 tx: request_tx,
711 prologue: None,
714 }))
715 .is_err()
716 {
717 return Err(error!(
718 "The requester of the request stream has been dropped before the connection was established"
719 ));
720 }
721
722 request_stream_send_handler(writer, request_rx, context).await;
723 Ok(())
724 }
725
726 async fn request_stream_send_handler(
736 mut framed_writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
737 mut request_rx: mpsc::Receiver<TwoPartMessage>,
738 context: Arc<dyn AsyncEngineContext>,
739 ) {
740 let killed = context.killed();
743 let stopped = context.stopped();
744 tokio::pin!(killed, stopped);
745
746 let closing_msg: Option<ControlMessage> = loop {
747 tokio::select! {
748 biased;
749
750 _ = &mut killed => {
751 tracing::trace!("context kill received in request-stream send handler");
752 break Some(ControlMessage::Kill);
753 }
754
755 _ = &mut stopped => {
756 tracing::trace!("context stop received in request-stream send handler");
757 break Some(ControlMessage::Stop);
758 }
759
760 msg = request_rx.recv() => {
761 match msg {
762 Some(msg) => {
763 if let Err(e) = framed_writer.send(msg).await {
764 tracing::trace!(
765 "failed to send request-stream frame to downstream: {:?}",
766 e
767 );
768 break None;
769 }
770 }
771 None => {
772 tracing::trace!("upstream request-stream sender closed; sending sentinel");
773 break Some(ControlMessage::Sentinel);
774 }
775 }
776 }
777 }
778 };
779
780 if let Some(ctrl) = closing_msg
781 && let Ok(bytes) = serde_json::to_vec(&ctrl)
782 && let Err(err) = framed_writer
783 .send(TwoPartMessage::from_header(bytes.into()))
784 .await
785 {
786 tracing::trace!(?err, ?ctrl, "request-stream closing-frame send failed");
787 }
788
789 let mut inner = framed_writer.into_inner();
790 if let Err(err) = inner.flush().await {
791 tracing::trace!(?err, "request-stream socket flush failed");
792 }
793 if let Err(err) = inner.shutdown().await {
794 tracing::trace!(?err, "request-stream socket shutdown failed");
795 }
796 }
797
798 async fn process_response_stream(
799 subject: String,
800 state: Arc<Mutex<State>>,
801 mut reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
802 writer: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
803 ) -> Result<()> {
804 let response_stream = TcpStreamServer::take_response_stream(&state, &subject).ok_or_else(|| {
805 error!("Subject not found: {}; upstream publisher specified a subject unknown to the downsteam subscriber", subject)
806 })?;
807
808 let RequestedRecvConnection {
810 context,
811 connection,
812 send_buffer_count,
813 } = response_stream;
814
815 let prologue = reader
818 .next()
819 .await
820 .ok_or(error!("Connection closed without a ControlMessge"))??;
821
822 let prologue = match prologue.into_message_type() {
824 TwoPartMessageType::HeaderOnly(header) => {
825 let prologue: ResponseStreamPrologue = serde_json::from_slice(&header)
826 .map_err(|e| error!("Failed to deserialize ControlMessage: {}", e))?;
827 prologue
828 }
829 _ => {
830 let msg = "malformed prologue: expected HeaderOnly ControlMessage";
835 let _ = connection.send(Err(msg.to_string()));
836 return Err(error!(msg));
837 }
838 };
839
840 if let Some(error) = &prologue.error {
847 let _ = connection.send(Err(error.clone()));
848 return Err(error!("Received error prologue: {}", error));
849 }
850
851 let (response_tx, response_rx) = data_plane_channel(send_buffer_count);
855
856 if connection
857 .send(Ok(crate::pipeline::network::StreamReceiver {
858 rx: response_rx,
859 }))
860 .is_err()
861 {
862 return Err(error!(
863 "The requester of the stream has been dropped before the connection was established"
864 ));
865 }
866
867 let (control_tx, control_rx) = mpsc::channel::<ControlMessage>(1);
868
869 let send_task = tokio::spawn(network_send_handler(writer, control_rx));
873
874 let recv_task = tokio::spawn(network_receive_handler(
876 reader,
877 response_tx,
878 control_tx,
879 context.clone(),
880 ));
881
882 let (monitor_result, forward_result) = tokio::join!(send_task, recv_task);
884
885 monitor_result?;
886 forward_result?;
887
888 Ok(())
889 }
890
891 async fn network_receive_handler(
892 mut framed_reader: FramedRead<tokio::io::ReadHalf<tokio::net::TcpStream>, TwoPartCodec>,
893 response_tx: mpsc::Sender<Bytes>,
894 control_tx: mpsc::Sender<ControlMessage>,
895 context: Arc<dyn AsyncEngineContext>,
896 ) {
897 let response_closed = response_tx.closed();
900 let killed = context.killed();
901 let stopped = context.stopped();
902 tokio::pin!(response_closed, killed, stopped);
903
904 let mut can_stop = true;
906 loop {
907 tokio::select! {
908 biased;
909
910 _ = &mut response_closed => {
911 tracing::trace!("response channel closed before the client finished writing data");
912 let _ = control_tx.send(ControlMessage::Kill).await;
913 break;
914 }
915
916 _ = &mut killed => {
917 tracing::trace!("context kill signal received; shutting down");
918 let _ = control_tx.send(ControlMessage::Kill).await;
919 break;
920 }
921
922 _ = &mut stopped, if can_stop => {
923 tracing::trace!("context stop signal received; shutting down");
924 can_stop = false;
927 let _ = control_tx.send(ControlMessage::Stop).await;
928 }
929
930 msg = framed_reader.next() => {
931 match msg {
932 Some(Ok(msg)) => {
933 let (header, data) = msg.into_parts();
934
935 if !header.is_empty() {
937 match process_control_message(header) {
938 Ok(ControlAction::Continue) => {}
939 Ok(ControlAction::Shutdown) => {
940 if !data.is_empty() {
941 tracing::warn!(
945 data_len = data.len(),
946 "client sent Sentinel with data (protocol violation); killing stream"
947 );
948 let _ = control_tx.send(ControlMessage::Kill).await;
949 break;
950 }
951 tracing::trace!("received sentinel message; shutting down");
952 break;
953 }
954 Err(e) => {
955 tracing::warn!(err = ?e, "malformed control message, closing connection");
958 let _ = control_tx.send(ControlMessage::Kill).await;
959 break;
960 }
961 }
962 }
963
964 if !data.is_empty()
965 && let Err(err) = response_tx.send(data).await {
966 tracing::debug!(?err, "forwarding body/data to response channel failed");
967 let _ = control_tx.send(ControlMessage::Kill).await;
968 break;
969 };
970 }
971 Some(Err(e)) => {
972 tracing::warn!(err = ?e, "tcp stream read error from worker, closing connection");
975 let _ = control_tx.send(ControlMessage::Kill).await;
976 break;
977 }
978 None => {
979 tracing::trace!("tcp stream was closed by client");
985 break;
986 }
987 }
988 }
989
990 }
991 }
992 }
993
994 async fn network_send_handler(
995 socket_tx: FramedWrite<tokio::io::WriteHalf<tokio::net::TcpStream>, TwoPartCodec>,
996 control_rx: mpsc::Receiver<ControlMessage>,
997 ) {
998 let mut socket_tx = socket_tx;
999 let mut control_rx = control_rx;
1000
1001 while let Some(control_msg) = control_rx.recv().await {
1002 if matches!(control_msg, ControlMessage::Sentinel) {
1006 tracing::warn!("received sentinel on send-side control channel; dropping");
1007 continue;
1008 }
1009 let bytes = match serde_json::to_vec(&control_msg) {
1010 Ok(b) => b,
1011 Err(e) => {
1012 tracing::warn!(err = ?e, ?control_msg, "failed to serialize control message");
1015 continue;
1016 }
1017 };
1018 let message = TwoPartMessage::from_header(bytes.into());
1019 match socket_tx.send(message).await {
1020 Ok(_) => tracing::debug!(?control_msg, "issued control message"),
1021 Err(e) => {
1022 tracing::debug!(err = ?e, ?control_msg, "failed to send control message")
1023 }
1024 }
1025 }
1026
1027 let mut inner = socket_tx.into_inner();
1028 if let Err(e) = inner.flush().await {
1029 tracing::debug!("failed to flush socket: {e}");
1030 }
1031 if let Err(e) = inner.shutdown().await {
1032 tracing::debug!("failed to shutdown socket: {e}");
1033 }
1034 }
1035}
1036
1037enum ControlAction {
1038 Continue,
1039 Shutdown,
1040}
1041
1042fn process_control_message(message: Bytes) -> Result<ControlAction> {
1043 match serde_json::from_slice::<ControlMessage>(&message)? {
1044 ControlMessage::Sentinel => {
1045 tracing::trace!("sentinel received; shutting down");
1048 Ok(ControlAction::Shutdown)
1049 }
1050 ControlMessage::Kill | ControlMessage::Stop => {
1051 anyhow::bail!("unexpected control message on response stream");
1055 }
1056 }
1057}
1058
1059#[cfg(test)]
1060mod tests {
1061 use super::*;
1062 use crate::engine::AsyncEngineContextProvider;
1063 use crate::pipeline::Context;
1064 use crate::pipeline::network::DEFAULT_SEND_BUFFER_COUNT;
1065 use crate::pipeline::network::tcp::client::TcpClient;
1066 use tokio::io::{AsyncWriteExt, ReadHalf, WriteHalf};
1067 use tokio::net::TcpStream;
1068
1069 struct FailingIpResolver;
1071
1072 impl IpResolver for FailingIpResolver {
1073 fn local_ip(&self) -> Result<std::net::IpAddr, Error> {
1074 Err(Error::LocalIpAddressNotFound)
1075 }
1076
1077 fn local_ipv6(&self) -> Result<std::net::IpAddr, Error> {
1078 Err(Error::LocalIpAddressNotFound)
1079 }
1080 }
1081
1082 #[tokio::test]
1083 async fn test_tcp_stream_server_default_behavior() {
1084 let options = ServerOptions::default();
1087 let result = TcpStreamServer::new(options).await;
1088
1089 assert!(
1090 result.is_ok(),
1091 "TcpStreamServer::new should succeed with default options"
1092 );
1093
1094 let server = result.unwrap();
1095
1096 let context = Context::new(());
1098 let stream_options = StreamOptions::builder()
1099 .context(context.context())
1100 .enable_request_stream(false)
1101 .enable_response_stream(true)
1102 .build()
1103 .unwrap();
1104
1105 let pending_connection = server.register(stream_options).await;
1106
1107 let connection_info = pending_connection
1109 .recv_stream
1110 .as_ref()
1111 .unwrap()
1112 .connection_info
1113 .clone();
1114
1115 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1116 let socket_addr = tcp_info.address.parse::<std::net::SocketAddr>().unwrap();
1117
1118 assert!(
1120 socket_addr.port() > 0,
1121 "Server should be assigned a valid port number"
1122 );
1123
1124 println!(
1125 "Server created successfully with address: {}",
1126 tcp_info.address
1127 );
1128 }
1129
1130 #[test]
1136 fn data_plane_channel_capacity_matches_send_buffer_count() {
1137 let (tx, _rx) = data_plane_channel::<()>(7);
1138 assert_eq!(tx.max_capacity(), 7);
1139
1140 let (tx, _rx) = data_plane_channel::<()>(DEFAULT_SEND_BUFFER_COUNT);
1141 assert_eq!(tx.max_capacity(), 64);
1142
1143 let (tx, _rx) = data_plane_channel::<()>(0);
1145 assert_eq!(tx.max_capacity(), 1);
1146 }
1147
1148 #[tokio::test]
1153 async fn register_threads_send_buffer_count_into_connection_structs() {
1154 let server = TcpStreamServer::new(ServerOptions::default())
1155 .await
1156 .expect("server");
1157 let context = Context::new(());
1158 let options = StreamOptions::builder()
1159 .context(context.context())
1160 .enable_request_stream(true)
1161 .enable_response_stream(true)
1162 .send_buffer_count(7)
1163 .build()
1164 .unwrap();
1165
1166 let _pending = server.register(options).await;
1167
1168 let state = server.state.lock();
1169 assert_eq!(state.tx_subjects.len(), 1, "one request stream registered");
1170 assert_eq!(state.rx_subjects.len(), 1, "one response stream registered");
1171 assert!(
1172 state.tx_subjects.values().all(|c| c.send_buffer_count == 7),
1173 "send_buffer_count must reach RequestedSendConnection"
1174 );
1175 assert!(
1176 state.rx_subjects.values().all(|c| c.send_buffer_count == 7),
1177 "send_buffer_count must reach RequestedRecvConnection"
1178 );
1179 }
1180
1181 #[tokio::test]
1182 async fn test_tcp_stream_server_fallback_to_loopback() {
1183 let options = ServerOptions::builder().port(0).build().unwrap();
1187
1188 let result = TcpStreamServer::new_with_resolver(options, FailingIpResolver).await;
1190 assert!(
1191 result.is_ok(),
1192 "Server creation should succeed with fallback even when IP detection fails"
1193 );
1194
1195 let server = result.unwrap();
1196
1197 let context = Context::new(());
1199 let stream_options = StreamOptions::builder()
1200 .context(context.context())
1201 .enable_request_stream(false)
1202 .enable_response_stream(true)
1203 .build()
1204 .unwrap();
1205
1206 let pending_connection = server.register(stream_options).await;
1207 let connection_info = pending_connection
1208 .recv_stream
1209 .as_ref()
1210 .unwrap()
1211 .connection_info
1212 .clone();
1213
1214 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1215 let socket_addr = tcp_info.address.parse::<std::net::SocketAddr>().unwrap();
1216
1217 let ip = socket_addr.ip();
1219 assert!(
1220 ip.is_loopback(),
1221 "Should use loopback when IP detection fails"
1222 );
1223
1224 assert_eq!(
1226 ip,
1227 std::net::IpAddr::V4(std::net::Ipv4Addr::new(127, 0, 0, 1)),
1228 "Fallback should use exactly 127.0.0.1, got: {}",
1229 ip
1230 );
1231
1232 println!("SUCCESS: Fallback to 127.0.0.1 was confirmed: {}", ip);
1233
1234 assert!(socket_addr.port() > 0, "Server should have a valid port");
1236 }
1237
1238 async fn test_server() -> Arc<TcpStreamServer> {
1240 TcpStreamServer::new_with_resolver(
1241 ServerOptions::builder().port(0).build().unwrap(),
1242 FailingIpResolver,
1243 )
1244 .await
1245 .unwrap()
1246 }
1247
1248 async fn register_and_get_subject(
1250 server: &TcpStreamServer,
1251 ) -> (
1252 String,
1253 tokio::sync::oneshot::Receiver<Result<super::StreamReceiver, String>>,
1254 ) {
1255 let context = Context::new(());
1256 let options = StreamOptions::builder()
1257 .context(context.context())
1258 .enable_request_stream(false)
1259 .enable_response_stream(true)
1260 .build()
1261 .unwrap();
1262
1263 let pending = server.register(options).await;
1264 let recv_stream = pending.recv_stream.unwrap();
1265 let (conn_info, provider) = recv_stream.into_parts();
1266 let tcp_info: TcpStreamConnectionInfo = conn_info.try_into().unwrap();
1267 (tcp_info.subject, provider)
1268 }
1269
1270 fn make_eid(
1272 namespace: &str,
1273 component: &str,
1274 endpoint: &str,
1275 instance_id: u64,
1276 ) -> EndpointInstanceId {
1277 EndpointInstanceId {
1278 namespace: namespace.to_string(),
1279 component: component.to_string(),
1280 endpoint: endpoint.to_string(),
1281 instance_id,
1282 }
1283 }
1284
1285 async fn register_and_get_bidi_subjects(
1288 server: &TcpStreamServer,
1289 ) -> (
1290 String,
1291 tokio::sync::oneshot::Receiver<Result<super::StreamSender, String>>,
1292 String,
1293 tokio::sync::oneshot::Receiver<Result<super::StreamReceiver, String>>,
1294 ) {
1295 let context = Context::new(());
1296 let options = StreamOptions::builder()
1297 .context(context.context())
1298 .enable_request_stream(true)
1299 .enable_response_stream(true)
1300 .build()
1301 .unwrap();
1302
1303 let pending = server.register(options).await;
1304 let send_stream = pending.send_stream.unwrap();
1305 let recv_stream = pending.recv_stream.unwrap();
1306 let (send_info, send_provider) = send_stream.into_parts();
1307 let (recv_info, recv_provider) = recv_stream.into_parts();
1308 let send_tcp_info: TcpStreamConnectionInfo = send_info.try_into().unwrap();
1309 let recv_tcp_info: TcpStreamConnectionInfo = recv_info.try_into().unwrap();
1310 (
1311 send_tcp_info.subject,
1312 send_provider,
1313 recv_tcp_info.subject,
1314 recv_provider,
1315 )
1316 }
1317
1318 #[tokio::test]
1323 async fn test_cancel_instance_streams_drops_both_bidi_halves() {
1324 let server = test_server().await;
1325 let (send_subj, send_provider, recv_subj, recv_provider) =
1326 register_and_get_bidi_subjects(&server).await;
1327
1328 let id = make_eid("ns", "comp", "generate", 7);
1329 assert!(
1330 server
1331 .associate_instance(&recv_subj, Some(&send_subj), &id)
1332 .await,
1333 "fresh instance must not be tombstoned"
1334 );
1335
1336 let cancelled = server.cancel_instance_streams(&id).await;
1337 assert_eq!(cancelled, 2, "both request + response halves must count");
1338
1339 assert!(
1340 recv_provider.await.is_err(),
1341 "recv provider should resolve with RecvError"
1342 );
1343 assert!(
1344 send_provider.await.is_err(),
1345 "send provider should resolve with RecvError after instance cancellation"
1346 );
1347 }
1348
1349 #[tokio::test]
1352 async fn test_associate_instance_tombstone_cancels_both_bidi_halves() {
1353 let server = test_server().await;
1354 let id = make_eid("ns", "comp", "generate", 8);
1355 server.cancel_instance_streams(&id).await;
1357
1358 let (send_subj, send_provider, recv_subj, recv_provider) =
1359 register_and_get_bidi_subjects(&server).await;
1360
1361 assert!(
1362 !server
1363 .associate_instance(&recv_subj, Some(&send_subj), &id)
1364 .await,
1365 "tombstoned instance must reject association"
1366 );
1367
1368 assert!(recv_provider.await.is_err());
1369 assert!(send_provider.await.is_err());
1370 }
1371
1372 #[tokio::test]
1373 async fn test_cancel_instance_streams_unblocks_receiver() {
1374 let server = test_server().await;
1375
1376 let (subject, provider) = register_and_get_subject(&server).await;
1377
1378 let id = make_eid("ns", "comp", "generate", 42);
1379 assert!(server.associate_instance(&subject, None, &id).await);
1380
1381 let cancelled = server.cancel_instance_streams(&id).await;
1382 assert_eq!(cancelled, 1);
1383
1384 let result = provider.await;
1386 assert!(result.is_err(), "Expected RecvError after cancellation");
1387 }
1388
1389 #[tokio::test]
1390 async fn test_cancel_instance_streams_multiple_subjects() {
1391 let server = test_server().await;
1392
1393 let (subj1, prov1) = register_and_get_subject(&server).await;
1394 let (subj2, prov2) = register_and_get_subject(&server).await;
1395 let (subj3, prov3) = register_and_get_subject(&server).await;
1396
1397 let id10 = make_eid("ns", "comp", "generate", 10);
1398 let id20 = make_eid("ns", "comp", "generate", 20);
1399
1400 assert!(server.associate_instance(&subj1, None, &id10).await);
1402 assert!(server.associate_instance(&subj2, None, &id10).await);
1403 assert!(server.associate_instance(&subj3, None, &id20).await);
1404
1405 let cancelled = server.cancel_instance_streams(&id10).await;
1407 assert_eq!(cancelled, 2);
1408
1409 assert!(prov1.await.is_err());
1410 assert!(prov2.await.is_err());
1411
1412 let cancelled = server.cancel_instance_streams(&id20).await;
1414 assert_eq!(cancelled, 1);
1415 assert!(prov3.await.is_err());
1416 }
1417
1418 #[tokio::test]
1419 async fn test_cancel_instance_streams_nonexistent_instance() {
1420 let server = test_server().await;
1421
1422 let id = make_eid("ns", "comp", "generate", 999);
1423 let cancelled = server.cancel_instance_streams(&id).await;
1424 assert_eq!(cancelled, 0);
1425 }
1426
1427 #[tokio::test]
1428 async fn test_cancel_recv_stream_cleans_up_instance_tracking() {
1429 let server = test_server().await;
1430
1431 let (subject, _provider) = register_and_get_subject(&server).await;
1432 let id = make_eid("ns", "comp", "generate", 42);
1433 assert!(server.associate_instance(&subject, None, &id).await);
1434
1435 server.cancel_recv_stream(&subject).await;
1437
1438 let cancelled = server.cancel_instance_streams(&id).await;
1440 assert_eq!(
1441 cancelled, 0,
1442 "Instance tracking should have been cleaned up"
1443 );
1444 }
1445
1446 #[tokio::test]
1447 async fn test_registered_stream_drop_runs_cleanup() {
1448 let server = test_server().await;
1449
1450 let context = Context::new(());
1452 let options = StreamOptions::builder()
1453 .context(context.context())
1454 .enable_request_stream(false)
1455 .enable_response_stream(true)
1456 .build()
1457 .unwrap();
1458
1459 let pending = server.register(options).await;
1460 let recv_stream = pending.recv_stream.unwrap();
1461
1462 let tcp_info: TcpStreamConnectionInfo =
1464 recv_stream.connection_info.clone().try_into().unwrap();
1465 let subject = tcp_info.subject.clone();
1466
1467 {
1469 let state = server.state.lock();
1470 assert!(state.rx_subjects.contains_key(&subject));
1471 }
1472
1473 drop(recv_stream);
1475
1476 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1478
1479 {
1481 let state = server.state.lock();
1482 assert!(
1483 !state.rx_subjects.contains_key(&subject),
1484 "RAII cleanup should have removed the rx_subjects entry"
1485 );
1486 }
1487 }
1488
1489 #[tokio::test]
1490 async fn test_registered_stream_into_parts_disarms_cleanup() {
1491 let server = test_server().await;
1492
1493 let context = Context::new(());
1494 let options = StreamOptions::builder()
1495 .context(context.context())
1496 .enable_request_stream(false)
1497 .enable_response_stream(true)
1498 .build()
1499 .unwrap();
1500
1501 let pending = server.register(options).await;
1502 let recv_stream = pending.recv_stream.unwrap();
1503
1504 let tcp_info: TcpStreamConnectionInfo =
1505 recv_stream.connection_info.clone().try_into().unwrap();
1506 let subject = tcp_info.subject.clone();
1507
1508 let (_conn_info, _provider) = recv_stream.into_parts();
1510
1511 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1513
1514 {
1516 let state = server.state.lock();
1517 assert!(
1518 state.rx_subjects.contains_key(&subject),
1519 "into_parts() should disarm the RAII cleanup"
1520 );
1521 }
1522 }
1523
1524 #[tokio::test]
1525 async fn test_associate_after_cancel_is_immediately_cancelled() {
1526 let server = test_server().await;
1528
1529 let id = make_eid("ns", "comp", "generate", 42);
1530
1531 let cancelled = server.cancel_instance_streams(&id).await;
1533 assert_eq!(cancelled, 0);
1534
1535 let (subject, provider) = register_and_get_subject(&server).await;
1537 let associated = server.associate_instance(&subject, None, &id).await;
1538
1539 assert!(
1541 !associated,
1542 "associate_instance on a tombstoned instance should return false"
1543 );
1544
1545 let result = provider.await;
1548 assert!(
1549 result.is_err(),
1550 "Late associate_instance on a tombstoned instance should immediately cancel"
1551 );
1552 }
1553
1554 #[tokio::test]
1555 async fn test_clear_tombstone_allows_new_associations() {
1556 let server = test_server().await;
1557
1558 let id = make_eid("ns", "comp", "generate", 42);
1559
1560 server.cancel_instance_streams(&id).await;
1561 server.clear_instance_tombstone(&id).await;
1562
1563 let (subject, _provider) = register_and_get_subject(&server).await;
1565 assert!(server.associate_instance(&subject, None, &id).await);
1566
1567 let cancelled = server.cancel_instance_streams(&id).await;
1569 assert_eq!(
1570 cancelled, 1,
1571 "After clearing tombstone, subjects should be tracked normally"
1572 );
1573 }
1574
1575 #[tokio::test]
1576 async fn test_cancel_does_not_affect_sibling_endpoint() {
1577 let server = test_server().await;
1580
1581 let (gen_subj, gen_prov) = register_and_get_subject(&server).await;
1582 let (pre_subj, pre_prov) = register_and_get_subject(&server).await;
1583
1584 let gen_id = make_eid("ns", "comp", "generate", 42);
1585 let pre_id = make_eid("ns", "comp", "prefill", 42);
1586
1587 assert!(server.associate_instance(&gen_subj, None, &gen_id).await);
1588 assert!(server.associate_instance(&pre_subj, None, &pre_id).await);
1589
1590 let cancelled = server.cancel_instance_streams(&gen_id).await;
1592 assert_eq!(
1593 cancelled, 1,
1594 "Only the generate subject should be cancelled"
1595 );
1596 assert!(gen_prov.await.is_err());
1597
1598 let still_pending = server.cancel_instance_streams(&pre_id).await;
1600 assert_eq!(still_pending, 1, "prefill subject should still be tracked");
1601 assert!(pre_prov.await.is_err());
1602 }
1603
1604 #[tokio::test]
1605 async fn test_tombstone_is_endpoint_scoped() {
1606 let server = test_server().await;
1609
1610 let gen_id = make_eid("ns", "comp", "generate", 42);
1611 let pre_id = make_eid("ns", "comp", "prefill", 42);
1612
1613 server.cancel_instance_streams(&gen_id).await;
1614
1615 let (gen_subj, gen_prov) = register_and_get_subject(&server).await;
1617 assert!(
1618 !server.associate_instance(&gen_subj, None, &gen_id).await,
1619 "generate should be tombstoned"
1620 );
1621 assert!(gen_prov.await.is_err());
1622
1623 let (pre_subj, _pre_prov) = register_and_get_subject(&server).await;
1625 assert!(
1626 server.associate_instance(&pre_subj, None, &pre_id).await,
1627 "prefill tombstone is independent; subject should be tracked"
1628 );
1629 let count = server.cancel_instance_streams(&pre_id).await;
1630 assert_eq!(count, 1, "prefill subject should be tracked normally");
1631 }
1632
1633 #[tokio::test]
1634 async fn test_cancel_does_not_affect_different_component() {
1635 let server = test_server().await;
1639
1640 let (subj_a, prov_a) = register_and_get_subject(&server).await;
1641 let (subj_b, prov_b) = register_and_get_subject(&server).await;
1642
1643 let id_a = make_eid("ns-a", "comp-a", "generate", 42);
1645 let id_b = make_eid("ns-b", "comp-b", "generate", 42);
1646
1647 assert!(server.associate_instance(&subj_a, None, &id_a).await);
1648 assert!(server.associate_instance(&subj_b, None, &id_b).await);
1649
1650 let cancelled = server.cancel_instance_streams(&id_a).await;
1652 assert_eq!(cancelled, 1, "Only service-A subject should be cancelled");
1653 assert!(prov_a.await.is_err());
1654
1655 let still_tracked = server.cancel_instance_streams(&id_b).await;
1657 assert_eq!(still_tracked, 1, "Service-B subject should be unaffected");
1658 assert!(prov_b.await.is_err());
1659 }
1660
1661 #[tokio::test(start_paused = true)]
1662 async fn test_tombstone_expires_after_ttl() {
1663 let server = test_server().await;
1667
1668 let id = make_eid("ns", "comp", "generate", 42);
1669
1670 server.cancel_instance_streams(&id).await;
1672 {
1673 let state = server.state.lock();
1674 assert!(state.removed_instances.contains_key(&id));
1675 }
1676
1677 tokio::time::advance(TOMBSTONE_TTL + Duration::from_secs(1)).await;
1679
1680 let (subject, _provider) = register_and_get_subject(&server).await;
1683 assert!(
1684 server.associate_instance(&subject, None, &id).await,
1685 "tombstone older than TTL should not block association"
1686 );
1687
1688 {
1691 let state = server.state.lock();
1692 assert!(
1693 !state.removed_instances.contains_key(&id),
1694 "expired tombstone should be pruned, not retained"
1695 );
1696 }
1697 }
1698
1699 #[tokio::test(start_paused = true)]
1700 async fn test_tombstone_within_ttl_blocks_associate() {
1701 let server = test_server().await;
1704
1705 let id = make_eid("ns", "comp", "generate", 42);
1706 server.cancel_instance_streams(&id).await;
1707
1708 tokio::time::advance(Duration::from_secs(1)).await;
1710
1711 let (subject, provider) = register_and_get_subject(&server).await;
1712 assert!(
1713 !server.associate_instance(&subject, None, &id).await,
1714 "tombstone within TTL must still block association"
1715 );
1716 assert!(provider.await.is_err());
1717 }
1718
1719 #[tokio::test(start_paused = true)]
1720 async fn test_tombstone_lazy_prune_on_cancel() {
1721 let server = test_server().await;
1724
1725 let id_old = make_eid("ns", "comp", "generate", 1);
1726 let id_new = make_eid("ns", "comp", "generate", 2);
1727
1728 server.cancel_instance_streams(&id_old).await;
1729 tokio::time::advance(TOMBSTONE_TTL + Duration::from_secs(1)).await;
1730 server.cancel_instance_streams(&id_new).await;
1731
1732 let state = server.state.lock();
1733 assert!(
1734 !state.removed_instances.contains_key(&id_old),
1735 "old tombstone should be pruned by the next cancel_instance_streams call"
1736 );
1737 assert!(
1738 state.removed_instances.contains_key(&id_new),
1739 "fresh tombstone should be retained"
1740 );
1741 assert_eq!(state.removed_instances.len(), 1);
1742 }
1743
1744 #[tokio::test]
1745 async fn test_clear_tombstone_only_affects_named_identity() {
1746 let server = test_server().await;
1751
1752 let id_a = make_eid("ns", "comp", "generate", 1);
1753 let id_b = make_eid("ns", "comp", "generate", 2);
1754
1755 server.cancel_instance_streams(&id_a).await;
1756 server.clear_instance_tombstone(&id_b).await;
1757
1758 let state = server.state.lock();
1759 assert!(
1760 state.removed_instances.contains_key(&id_a),
1761 "clearing a different identity must not remove id_a's tombstone"
1762 );
1763 }
1764
1765 #[tokio::test]
1766 async fn test_tombstone_scoped_to_full_identity() {
1767 let server = test_server().await;
1770
1771 let id_a = make_eid("ns-a", "comp-a", "generate", 42);
1772 let id_b = make_eid("ns-b", "comp-b", "generate", 42);
1773
1774 server.cancel_instance_streams(&id_a).await;
1776
1777 let (subj_a, prov_a) = register_and_get_subject(&server).await;
1779 assert!(!server.associate_instance(&subj_a, None, &id_a).await);
1780 assert!(prov_a.await.is_err());
1781
1782 let (subj_b, _prov_b) = register_and_get_subject(&server).await;
1784 assert!(
1785 server.associate_instance(&subj_b, None, &id_b).await,
1786 "Different namespace/component must not be tombstoned"
1787 );
1788 assert_eq!(server.cancel_instance_streams(&id_b).await, 1);
1789 }
1790
1791 type TestFramedRead = FramedRead<ReadHalf<TcpStream>, TwoPartCodec>;
1792 type TestFramedWrite = FramedWrite<WriteHalf<TcpStream>, TwoPartCodec>;
1793 type TestResponseStream = (TestFramedRead, TestFramedWrite, StreamReceiver);
1794
1795 async fn open_registered_response_stream() -> TestResponseStream {
1799 let options = ServerOptions::builder().port(0).build().unwrap();
1800 let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver)
1801 .await
1802 .unwrap();
1803 let context = Context::new(());
1804 let stream_options = StreamOptions::builder()
1805 .context(context.context())
1806 .enable_request_stream(false)
1807 .enable_response_stream(true)
1808 .build()
1809 .unwrap();
1810 let pending_connection = server.register(stream_options).await;
1811 let registered_stream = pending_connection.recv_stream.unwrap();
1812 let (connection_info, stream_provider) = registered_stream.into_parts();
1813 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1814
1815 let stream = TcpStream::connect(&tcp_info.address).await.unwrap();
1816 let (read_half, write_half) = tokio::io::split(stream);
1817 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
1818 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
1819
1820 let handshake = CallHomeHandshake {
1821 subject: tcp_info.subject,
1822 stream_type: StreamType::Response,
1823 };
1824 framed_writer
1825 .send(TwoPartMessage::from_header(
1826 serde_json::to_vec(&handshake).unwrap().into(),
1827 ))
1828 .await
1829 .unwrap();
1830 framed_writer
1831 .send(TwoPartMessage::from_header(
1832 serde_json::to_vec(&ResponseStreamPrologue { error: None })
1833 .unwrap()
1834 .into(),
1835 ))
1836 .await
1837 .unwrap();
1838
1839 let receiver = tokio::time::timeout(std::time::Duration::from_secs(1), stream_provider)
1842 .await
1843 .expect("server should establish response stream within timeout")
1844 .expect("stream provider should not be dropped")
1845 .expect("response stream should be accepted");
1846
1847 (framed_reader, framed_writer, receiver)
1848 }
1849
1850 async fn recv_control_message(framed_reader: &mut TestFramedRead) -> ControlMessage {
1851 let message = tokio::time::timeout(std::time::Duration::from_secs(1), framed_reader.next())
1854 .await
1855 .expect("server should send a control message within timeout")
1856 .expect("server should not close before sending control")
1857 .expect("control message should decode");
1858 let (header, data) = message.optional_parts();
1859 assert!(data.is_none(), "control message should not contain data");
1860 serde_json::from_slice(header.expect("control header missing").as_ref()).unwrap()
1861 }
1862
1863 #[tokio::test]
1868 async fn test_tcp_stream_server_sends_kill_on_unexpected_control_message() {
1869 let (mut framed_reader, mut framed_writer, _receiver) =
1870 open_registered_response_stream().await;
1871
1872 framed_writer
1873 .send(TwoPartMessage::from_header(
1874 serde_json::to_vec(&ControlMessage::Stop).unwrap().into(),
1875 ))
1876 .await
1877 .unwrap();
1878
1879 assert_eq!(
1880 recv_control_message(&mut framed_reader).await,
1881 ControlMessage::Kill,
1882 "unexpected control message should kill only this stream"
1883 );
1884 }
1885
1886 #[tokio::test]
1890 async fn test_tcp_stream_server_sends_kill_on_read_error() {
1891 let (mut framed_reader, framed_writer, _receiver) = open_registered_response_stream().await;
1892
1893 let mut raw_writer = framed_writer.into_inner();
1894 raw_writer.write_all(&[0u8; 8]).await.unwrap();
1895 raw_writer.shutdown().await.unwrap();
1896
1897 assert_eq!(
1898 recv_control_message(&mut framed_reader).await,
1899 ControlMessage::Kill,
1900 "framing read error should kill only this stream"
1901 );
1902 }
1903
1904 #[tokio::test]
1907 async fn test_tcp_stream_server_sends_kill_on_sentinel_with_data() {
1908 let (mut framed_reader, mut framed_writer, _receiver) =
1909 open_registered_response_stream().await;
1910
1911 let header = serde_json::to_vec(&ControlMessage::Sentinel)
1912 .unwrap()
1913 .into();
1914 framed_writer
1915 .send(TwoPartMessage::from_parts(
1916 header,
1917 Bytes::from_static(b"unexpected payload"),
1918 ))
1919 .await
1920 .unwrap();
1921
1922 assert_eq!(
1923 recv_control_message(&mut framed_reader).await,
1924 ControlMessage::Kill,
1925 "Sentinel with data should kill only this stream"
1926 );
1927 }
1928
1929 #[tokio::test]
1933 async fn test_tcp_stream_server_returns_error_on_invalid_prologue() {
1934 let options = ServerOptions::builder().port(0).build().unwrap();
1935 let server = TcpStreamServer::new_with_resolver(options, FailingIpResolver)
1936 .await
1937 .unwrap();
1938 let context = Context::new(());
1939 let stream_options = StreamOptions::builder()
1940 .context(context.context())
1941 .enable_request_stream(false)
1942 .enable_response_stream(true)
1943 .build()
1944 .unwrap();
1945 let pending_connection = server.register(stream_options).await;
1946 let registered_stream = pending_connection.recv_stream.unwrap();
1947 let (connection_info, stream_provider) = registered_stream.into_parts();
1948 let tcp_info: TcpStreamConnectionInfo = connection_info.try_into().unwrap();
1949
1950 let stream = TcpStream::connect(&tcp_info.address).await.unwrap();
1951 let (_read_half, write_half) = tokio::io::split(stream);
1952 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
1953
1954 let handshake = CallHomeHandshake {
1955 subject: tcp_info.subject,
1956 stream_type: StreamType::Response,
1957 };
1958 framed_writer
1959 .send(TwoPartMessage::from_header(
1960 serde_json::to_vec(&handshake).unwrap().into(),
1961 ))
1962 .await
1963 .unwrap();
1964
1965 framed_writer
1967 .send(TwoPartMessage::from_data(Bytes::from_static(
1968 b"not a prologue",
1969 )))
1970 .await
1971 .unwrap();
1972
1973 let outcome = tokio::time::timeout(std::time::Duration::from_secs(1), stream_provider)
1974 .await
1975 .expect("stream provider should resolve quickly")
1976 .expect("stream provider channel should not be dropped");
1977 match outcome {
1979 Err(err) => assert!(
1980 err.contains("malformed prologue"),
1981 "expected malformed-prologue error, got: {err}"
1982 ),
1983 Ok(_) => panic!("invalid prologue should produce an error, but got Ok"),
1984 }
1985 }
1986
1987 #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
1988 async fn test_concurrent_response_registration_and_call_home() {
1989 const STREAMS: usize = 128;
1990
1991 let result = time::timeout(Duration::from_secs(20), async {
1992 let server = test_server().await;
1993 let mut pending_streams = Vec::with_capacity(STREAMS);
1994 let mut client_tasks = Vec::with_capacity(STREAMS);
1995
1996 for idx in 0..STREAMS {
1997 let context = Context::new(());
1998 let options = StreamOptions::builder()
1999 .context(context.context())
2000 .enable_request_stream(false)
2001 .enable_response_stream(true)
2002 .build()
2003 .unwrap();
2004
2005 let pending = server.register(options).await;
2006 let registered_stream = pending.recv_stream.unwrap();
2007 let (connection_info, stream_provider) = registered_stream.into_parts();
2008 let client_context =
2009 Context::with_id_and_metadata((), context.id().to_string(), Default::default());
2010 let payload = Bytes::from(format!("payload-{idx}"));
2011
2012 pending_streams.push((idx, payload.clone(), stream_provider));
2013 client_tasks.push(tokio::spawn(async move {
2014 let mut sender = TcpClient::create_response_stream(
2015 client_context.context(),
2016 connection_info,
2017 None,
2018 )
2019 .await
2020 .unwrap();
2021 sender.send_prologue(None).await.unwrap();
2022 sender.send(payload).await.unwrap();
2023 }));
2024 }
2025
2026 for task in client_tasks {
2027 task.await.unwrap();
2028 }
2029
2030 for (idx, expected, stream_provider) in pending_streams {
2031 let mut stream = stream_provider.await.unwrap().unwrap();
2032 let actual = stream.rx.recv().await.unwrap();
2033 assert_eq!(actual, expected, "payload mismatch for stream {idx}");
2034 }
2035 })
2036 .await;
2037
2038 assert!(
2039 result.is_ok(),
2040 "concurrent response registration and call-home timed out"
2041 );
2042 }
2043
2044 use futures::SinkExt;
2052
2053 async fn register_and_dial_request_stream(
2058 server: &TcpStreamServer,
2059 ) -> (
2060 FramedRead<tokio::io::ReadHalf<TcpStream>, TwoPartCodec>,
2061 super::StreamSender,
2062 Arc<dyn AsyncEngineContext>,
2063 ) {
2064 let upstream_ctx = Context::new(()).context();
2065 let options = StreamOptions::builder()
2066 .context(upstream_ctx.clone())
2067 .enable_request_stream(true)
2068 .enable_response_stream(false)
2069 .build()
2070 .unwrap();
2071
2072 let pending = server.register(options).await;
2073 let send_stream = pending.send_stream.unwrap();
2074 let (conn_info, send_provider) = send_stream.into_parts();
2075 let tcp_info: TcpStreamConnectionInfo = conn_info.try_into().unwrap();
2076
2077 let raw = TcpStream::connect(&tcp_info.address).await.unwrap();
2078 let (read_half, write_half) = tokio::io::split(raw);
2079 let framed_reader = FramedRead::new(read_half, TwoPartCodec::default());
2080 let mut framed_writer = FramedWrite::new(write_half, TwoPartCodec::default());
2081
2082 let handshake = super::CallHomeHandshake {
2083 subject: tcp_info.subject.clone(),
2084 stream_type: StreamType::Request,
2085 };
2086 let handshake_bytes = serde_json::to_vec(&handshake).unwrap();
2087 framed_writer
2088 .send(TwoPartMessage::from_header(handshake_bytes.into()))
2089 .await
2090 .unwrap();
2091 drop(framed_writer);
2092
2093 let sender = send_provider.await.unwrap().unwrap();
2094 (framed_reader, sender, upstream_ctx)
2095 }
2096
2097 async fn next_control_message(
2100 reader: &mut FramedRead<tokio::io::ReadHalf<TcpStream>, TwoPartCodec>,
2101 ) -> ControlMessage {
2102 loop {
2103 let frame = reader
2104 .next()
2105 .await
2106 .expect("socket closed before control message arrived")
2107 .expect("decode error");
2108 if let Some(header) = frame.header() {
2109 return serde_json::from_slice::<ControlMessage>(header)
2110 .expect("invalid control message bytes");
2111 }
2112 }
2114 }
2115
2116 #[tokio::test]
2119 async fn test_request_stream_sends_sentinel_on_clean_drop() {
2120 let server = test_server().await;
2121 let (mut reader, sender, _ctx) = register_and_dial_request_stream(&server).await;
2122
2123 drop(sender);
2124
2125 let ctrl = next_control_message(&mut reader).await;
2126 assert!(
2127 matches!(ctrl, ControlMessage::Sentinel),
2128 "clean drain should emit Sentinel, got {ctrl:?}"
2129 );
2130 }
2131
2132 #[tokio::test]
2135 async fn test_request_stream_sends_kill_on_context_killed() {
2136 let server = test_server().await;
2137 let (mut reader, _sender, ctx) = register_and_dial_request_stream(&server).await;
2138
2139 ctx.kill();
2140
2141 let ctrl = next_control_message(&mut reader).await;
2142 assert!(
2143 matches!(ctrl, ControlMessage::Kill),
2144 "context.kill() should emit Kill, got {ctrl:?}"
2145 );
2146 }
2147
2148 #[tokio::test]
2151 async fn test_request_stream_sends_stop_on_context_stopped() {
2152 let server = test_server().await;
2153 let (mut reader, _sender, ctx) = register_and_dial_request_stream(&server).await;
2154
2155 ctx.stop();
2156
2157 let ctrl = next_control_message(&mut reader).await;
2158 assert!(
2159 matches!(ctrl, ControlMessage::Stop),
2160 "context.stop() should emit Stop, got {ctrl:?}"
2161 );
2162 }
2163}