1#[cfg(unix)]
2use std::os::unix::fs::PermissionsExt;
3use std::{
4 convert::TryInto,
5 fs,
6 marker::{Send, Sync},
7 net::SocketAddr,
8 path::PathBuf,
9 sync::{
10 Arc, Mutex,
11 atomic::{AtomicUsize, Ordering},
12 },
13 time::Duration,
14};
15
16use bytes::{Buf, Bytes, BytesMut};
17use futures::{
18 executor::block_on,
19 future::{self, OptionFuture},
20 sink::{Sink, SinkExt},
21 stream::{self, StreamExt, TryStreamExt},
22};
23use futures_util::{Future, FutureExt, future::BoxFuture};
24use ipnet::IpNet;
25use listenfd::ListenFd;
26use tokio::{
27 self,
28 io::{AsyncRead, AsyncWrite},
29 net::{TcpStream, UnixListener},
30 task::JoinHandle,
31 time::sleep,
32};
33use tokio_stream::wrappers::UnixListenerStream;
34use tokio_util::codec::{Framed, length_delimited};
35use tracing::{Instrument, Span, field};
36use vector_lib::{
37 lookup::OwnedValuePath,
38 tcp::TcpKeepaliveConfig,
39 tls::{CertificateMetadata, MaybeTlsIncomingStream, MaybeTlsSettings},
40};
41
42use super::net::{RequestLimiter, SocketListenAddr};
43use crate::{
44 SourceSender,
45 event::Event,
46 internal_events::{
47 ConnectionOpen, OpenGauge, SocketBindError, SocketMode, SocketReceiveError,
48 TcpBytesReceived, TcpSocketError, TcpSocketTlsConnectionError, UnixSocketError,
49 UnixSocketFileDeleteError,
50 },
51 shutdown::ShutdownSignal,
52 sources::{
53 Source,
54 util::{
55 AfterReadExt,
56 net::{MAX_IN_FLIGHT_EVENTS_TARGET, try_bind_tcp_listener},
57 },
58 },
59};
60
61const FSTRM_CONTROL_FRAME_LENGTH_MAX: usize = 512;
62const FSTRM_CONTROL_FIELD_CONTENT_TYPE_LENGTH_MAX: usize = 256;
63
64const PERMIT_HOLD_TIMEOUT_MS: u64 = 10;
69
70pub type FrameStreamSink = Box<dyn Sink<Bytes, Error = std::io::Error> + Send + Unpin>;
71
72pub struct FrameStreamReader {
73 response_sink: Mutex<FrameStreamSink>,
74 expected_content_type: String,
75 state: FrameStreamState,
76}
77
78struct FrameStreamState {
79 expect_control_frame: bool,
80 control_state: ControlState,
81 is_bidirectional: bool,
82}
83impl FrameStreamState {
84 const fn new() -> Self {
85 FrameStreamState {
86 expect_control_frame: false,
87 control_state: ControlState::Initial,
89 is_bidirectional: true, }
91 }
92}
93
94#[derive(PartialEq, Debug)]
95enum ControlState {
96 Initial,
97 GotReady,
98 ReadingData,
99 Stopped,
100}
101
102#[derive(Copy, Clone)]
103enum ControlHeader {
104 Accept,
105 Start,
106 Stop,
107 Ready,
108 Finish,
109}
110
111impl ControlHeader {
112 fn from_u32(val: u32) -> Result<Self, ()> {
113 match val {
114 0x01 => Ok(ControlHeader::Accept),
115 0x02 => Ok(ControlHeader::Start),
116 0x03 => Ok(ControlHeader::Stop),
117 0x04 => Ok(ControlHeader::Ready),
118 0x05 => Ok(ControlHeader::Finish),
119 _ => {
120 error!("Don't know header value {} (expected 0x01 - 0x05).", val);
121 Err(())
122 }
123 }
124 }
125
126 const fn to_u32(self) -> u32 {
127 match self {
128 ControlHeader::Accept => 0x01,
129 ControlHeader::Start => 0x02,
130 ControlHeader::Stop => 0x03,
131 ControlHeader::Ready => 0x04,
132 ControlHeader::Finish => 0x05,
133 }
134 }
135}
136
137enum ControlField {
138 ContentType,
139}
140
141impl ControlField {
142 fn from_u32(val: u32) -> Result<Self, ()> {
143 match val {
144 0x01 => Ok(ControlField::ContentType),
145 _ => {
146 error!("Don't know field type {} (expected 0x01).", val);
147 Err(())
148 }
149 }
150 }
151 const fn to_u32(&self) -> u32 {
152 match self {
153 ControlField::ContentType => 0x01,
154 }
155 }
156}
157
158fn advance_u32(b: &mut Bytes) -> Result<u32, ()> {
159 if b.len() < 4 {
160 error!("Malformed frame.");
161 return Err(());
162 }
163 let a = b.split_to(4);
164 Ok(u32::from_be_bytes(a[..].try_into().unwrap()))
165}
166
167impl FrameStreamReader {
168 pub fn new(response_sink: FrameStreamSink, expected_content_type: String) -> Self {
169 FrameStreamReader {
170 response_sink: Mutex::new(response_sink),
171 expected_content_type,
172 state: FrameStreamState::new(),
173 }
174 }
175
176 pub fn handle_frame(&mut self, frame: Bytes) -> Option<Bytes> {
177 if frame.is_empty() {
178 self.state.expect_control_frame = true;
180 None
181 } else if self.state.expect_control_frame {
182 self.state.expect_control_frame = false;
183 _ = self.handle_control_frame(frame);
184 None
185 } else {
186 if self.state.control_state == ControlState::ReadingData {
188 Some(frame) } else {
190 error!(
191 "Received a data frame while in state {:?}.",
192 self.state.control_state
193 );
194 None
195 }
196 }
197 }
198
199 fn handle_control_frame(&mut self, mut frame: Bytes) -> Result<(), ()> {
200 if frame.len() > FSTRM_CONTROL_FRAME_LENGTH_MAX {
202 error!("Control frame is too long.");
203 }
204
205 let header = ControlHeader::from_u32(advance_u32(&mut frame)?)?;
206
207 match self.state.control_state {
209 ControlState::Initial => {
210 match header {
211 ControlHeader::Ready => {
212 let content_type = self.process_fields(header, &mut frame)?.unwrap();
213
214 self.send_control_frame(Self::make_frame(
215 ControlHeader::Accept,
216 Some(content_type),
217 ));
218 self.state.control_state = ControlState::GotReady; }
220 ControlHeader::Start => {
221 _ = self.process_fields(header, &mut frame)?;
223 self.state.control_state = ControlState::ReadingData;
225 self.state.is_bidirectional = false; }
227 _ => error!("Got wrong control frame, expected READY."),
228 }
229 }
230 ControlState::GotReady => {
231 match header {
232 ControlHeader::Start => {
233 _ = self.process_fields(header, &mut frame)?;
235 self.state.control_state = ControlState::ReadingData;
237 }
238 _ => error!("Got wrong control frame, expected START."),
239 }
240 }
241 ControlState::ReadingData => {
242 match header {
243 ControlHeader::Stop => {
244 _ = self.process_fields(header, &mut frame)?;
246 if self.state.is_bidirectional {
247 self.send_control_frame(Self::make_frame(ControlHeader::Finish, None));
249 }
250 self.state.control_state = ControlState::Stopped; }
252 _ => error!("Got wrong control frame, expected STOP."),
253 }
254 }
255 ControlState::Stopped => error!("Unexpected control frame, current state is STOPPED."),
256 };
257 Ok(())
258 }
259
260 fn process_fields(
261 &mut self,
262 header: ControlHeader,
263 frame: &mut Bytes,
264 ) -> Result<Option<String>, ()> {
265 match header {
266 ControlHeader::Ready => {
267 let is_start_frame = false;
270 let content_type = self.process_content_type(frame, is_start_frame)?;
271 Ok(Some(content_type))
272 }
273 ControlHeader::Start => {
274 if frame.is_empty() {
276 Ok(None)
277 } else {
278 let is_start_frame = true;
280 let content_type = self.process_content_type(frame, is_start_frame)?;
281 Ok(Some(content_type))
282 }
283 }
284 ControlHeader::Stop => {
285 if !frame.is_empty() {
287 error!("Unexpected fields in STOP header.");
288 Err(())
289 } else {
290 Ok(None)
291 }
292 }
293 _ => {
294 error!("Unexpected control header value {:?}.", header.to_u32());
295 Err(())
296 }
297 }
298 }
299
300 fn process_content_type(&self, frame: &mut Bytes, is_start_frame: bool) -> Result<String, ()> {
301 if frame.is_empty() {
302 error!("No fields in control frame.");
303 return Err(());
304 }
305
306 let mut content_types = vec![];
307 while !frame.is_empty() {
308 let field_val = advance_u32(frame)?;
310 let field_type = ControlField::from_u32(field_val)?;
311 match field_type {
312 ControlField::ContentType => {
313 let field_len = advance_u32(frame)? as usize;
315
316 if field_len > FSTRM_CONTROL_FIELD_CONTENT_TYPE_LENGTH_MAX {
318 error!("Content-Type string is too long.");
319 return Err(());
320 }
321
322 let content_type = std::str::from_utf8(&frame[..field_len]).unwrap();
323 content_types.push(content_type.to_string());
324 frame.advance(field_len);
325 }
326 }
327 }
328
329 if is_start_frame && content_types.len() > 1 {
330 error!(
331 "START control frame can only have one content-type provided (got {}).",
332 content_types.len()
333 );
334 return Err(());
335 }
336
337 for content_type in &content_types {
338 if *content_type == self.expected_content_type {
339 return Ok(content_type.clone());
340 }
341 }
342
343 error!(
344 "Content types did not match up. Expected {} got {:?}.",
345 self.expected_content_type, content_types
346 );
347 Err(())
348 }
349
350 fn make_frame(header: ControlHeader, content_type: Option<String>) -> Bytes {
351 let mut frame = BytesMut::new();
352 frame.extend(header.to_u32().to_be_bytes());
353 if let Some(s) = content_type {
354 frame.extend(ControlField::ContentType.to_u32().to_be_bytes()); frame.extend((s.len() as u32).to_be_bytes()); frame.extend(s.as_bytes());
357 }
358 Bytes::from(frame)
359 }
360
361 fn send_control_frame(&mut self, frame: Bytes) {
362 let empty_frame = Bytes::from(&b""[..]); let mut stream = stream::iter(vec![Ok(empty_frame), Ok(frame)]);
364
365 if let Err(e) = block_on(self.response_sink.lock().unwrap().send_all(&mut stream)) {
366 error!("Encountered error '{:#?}' while sending control frame.", e);
367 }
368 }
369}
370
371pub trait FrameHandler {
372 fn content_type(&self) -> String;
373 fn max_frame_length(&self) -> usize;
374 fn handle_event(&self, received_from: Option<Bytes>, frame: Bytes) -> Option<Event>;
375 fn multithreaded(&self) -> bool;
376 fn max_frame_handling_tasks(&self) -> usize;
377 fn host_key(&self) -> &Option<OwnedValuePath>;
378 fn timestamp_key(&self) -> Option<&OwnedValuePath>;
379 fn source_type_key(&self) -> Option<&OwnedValuePath>;
380}
381
382pub trait UnixFrameHandler: FrameHandler {
383 fn socket_path(&self) -> PathBuf;
384 fn socket_file_mode(&self) -> Option<u32>;
385 fn socket_receive_buffer_size(&self) -> Option<usize>;
386 fn socket_send_buffer_size(&self) -> Option<usize>;
387}
388
389pub trait TcpFrameHandler: FrameHandler {
390 fn address(&self) -> SocketListenAddr;
391 fn keepalive(&self) -> Option<TcpKeepaliveConfig>;
392 fn shutdown_timeout_secs(&self) -> Duration;
393 fn tls(&self) -> MaybeTlsSettings;
394 fn tls_client_metadata_key(&self) -> Option<OwnedValuePath>;
395 fn receive_buffer_bytes(&self) -> Option<usize>;
396 fn max_connection_duration_secs(&self) -> Option<u64>;
397 fn max_connections(&self) -> Option<u32>;
398 fn allowed_origins(&self) -> Option<&[IpNet]>;
399 fn insert_tls_client_metadata(&mut self, metadata: Option<CertificateMetadata>);
400}
401
402pub fn build_framestream_tcp_source(
407 frame_handler: impl TcpFrameHandler + Send + Sync + Clone + 'static,
408 shutdown: ShutdownSignal,
409 out: SourceSender,
410) -> crate::Result<Source> {
411 let addr = frame_handler.address();
412 let tls = frame_handler.tls();
413 let shutdown = shutdown.clone();
414 let out = out.clone();
415
416 Ok(Box::pin(async move {
417 let listenfd = ListenFd::from_env();
418 let listener = try_bind_tcp_listener(
419 addr,
420 listenfd,
421 &tls,
422 None, frame_handler
424 .allowed_origins()
425 .map(|origins| origins.to_vec()),
426 )
427 .await
428 .map_err(|error| {
429 emit!(SocketBindError {
430 mode: SocketMode::Tcp,
431 error: &error,
432 })
433 })?;
434
435 info!(
436 message = "Listening.",
437 addr = %listener
438 .local_addr()
439 .map(SocketListenAddr::SocketAddr)
440 .unwrap_or(addr)
441 );
442
443 let tripwire = shutdown.clone();
444 let shutdown_timeout_secs = frame_handler.shutdown_timeout_secs();
445 let tripwire = async move {
446 _ = tripwire.await;
447 sleep(shutdown_timeout_secs).await;
448 }
449 .shared();
450
451 let connection_gauge = OpenGauge::new();
452 let shutdown_clone = shutdown.clone();
453
454 let request_limiter = RequestLimiter::new(
455 MAX_IN_FLIGHT_EVENTS_TARGET,
456 frame_handler.max_frame_handling_tasks(),
457 );
458
459 listener
460 .accept_stream_limited(frame_handler.max_connections())
461 .take_until(shutdown_clone)
462 .for_each(move |(connection, tcp_connection_permit)| {
463 let shutdown_signal = shutdown.clone();
464 let tripwire = tripwire.clone();
465 let out = out.clone();
466 let connection_gauge = connection_gauge.clone();
467 let request_limiter = request_limiter.clone();
468 let frame_handler_clone = frame_handler.clone();
469
470 async move {
471 let socket = match connection {
472 Ok(socket) => socket,
473 Err(error) => {
474 emit!(SocketReceiveError {
475 mode: SocketMode::Tcp,
476 error: &error
477 });
478 return;
479 }
480 };
481
482 let peer_addr = socket.peer_addr();
483 let span = info_span!("connection", %peer_addr);
484
485 let tripwire = tripwire
486 .map(move |_| {
487 info!(
488 message = "Resetting connection (still open after seconds).",
489 seconds = ?shutdown_timeout_secs
490 );
491 })
492 .boxed();
493
494 span.clone().in_scope(|| {
495 debug!(message = "Accepted a new connection.", peer_addr = %peer_addr);
496
497 let open_token =
498 connection_gauge.open(|count| emit!(ConnectionOpen { count }));
499
500 let fut = handle_stream(
501 frame_handler_clone,
502 shutdown_signal,
503 socket,
504 tripwire,
505 peer_addr,
506 out,
507 request_limiter,
508 );
509
510 tokio::spawn(
511 fut.map(move |()| {
512 drop(open_token);
513 drop(tcp_connection_permit);
514 })
515 .instrument(span.or_current()),
516 );
517 });
518 }
519 })
520 .map(Ok)
521 .await
522 }))
523}
524
525#[allow(clippy::too_many_arguments)]
526async fn handle_stream(
527 mut frame_handler: impl TcpFrameHandler + Send + Sync + Clone + 'static,
528 mut shutdown_signal: ShutdownSignal,
529 mut socket: MaybeTlsIncomingStream<TcpStream>,
530 mut tripwire: BoxFuture<'static, ()>,
531 peer_addr: SocketAddr,
532 out: SourceSender,
533 request_limiter: RequestLimiter,
534) {
535 tokio::select! {
536 result = socket.handshake() => {
537 if let Err(error) = result {
538 emit!(TcpSocketTlsConnectionError { error });
539 return;
540 }
541 },
542 _ = &mut shutdown_signal => {
543 return;
544 }
545 };
546
547 if let Some(keepalive) = frame_handler.keepalive()
548 && let Err(error) = socket.set_keepalive(keepalive)
549 {
550 warn!(message = "Failed configuring TCP keepalive.", %error);
551 }
552
553 if let Some(receive_buffer_bytes) = frame_handler.receive_buffer_bytes()
554 && let Err(error) = socket.set_receive_buffer_bytes(receive_buffer_bytes)
555 {
556 warn!(message = "Failed configuring receive buffer size on TCP socket.", %error);
557 }
558
559 let socket = socket.after_read(move |byte_size| {
560 emit!(TcpBytesReceived {
561 byte_size,
562 peer_addr
563 });
564 });
565
566 let certificate_metadata = socket
567 .get_ref()
568 .ssl_stream()
569 .and_then(|stream| stream.ssl().peer_certificate())
570 .map(CertificateMetadata::from);
571
572 frame_handler.insert_tls_client_metadata(certificate_metadata);
573
574 let span = info_span!("connection");
575 span.record("peer_addr", field::debug(&peer_addr));
576 let received_from: Option<Bytes> = Some(peer_addr.to_string().into());
577
578 let connection_close_timeout = OptionFuture::from(
579 frame_handler
580 .max_connection_duration_secs()
581 .map(|timeout_secs| tokio::time::sleep(Duration::from_secs(timeout_secs))),
582 );
583 tokio::pin!(connection_close_timeout);
584
585 let content_type = frame_handler.content_type();
586 let mut event_sink = out.clone();
587 let (sock_sink, sock_stream) = Framed::new(
588 socket,
589 length_delimited::Builder::new()
590 .max_frame_length(frame_handler.max_frame_length())
591 .new_codec(),
592 )
593 .split();
594 let mut reader = FrameStreamReader::new(Box::new(sock_sink), content_type);
595 let mut frames = sock_stream
596 .map_err(move |error| {
597 emit!(TcpSocketError {
598 error: &error,
599 peer_addr,
600 });
601 })
602 .filter_map(move |frame| {
603 future::ready(match frame {
604 Ok(f) => reader.handle_frame(Bytes::from(f)),
605 Err(_) => None,
606 })
607 });
608
609 let active_parsing_task_nums = Arc::new(AtomicUsize::new(0));
610 loop {
611 let mut permit = tokio::select! {
612 _ = &mut tripwire => break,
613 Some(_) = &mut connection_close_timeout => {
614 break;
615 },
616 _ = &mut shutdown_signal => {
617 break;
618 },
619 permit = request_limiter.acquire() => {
620 Some(permit)
621 }
622 else => break,
623 };
624
625 let timeout = tokio::time::sleep(Duration::from_millis(PERMIT_HOLD_TIMEOUT_MS));
626 tokio::pin!(timeout);
627
628 tokio::select! {
629 _ = &mut tripwire => break,
630 _ = &mut shutdown_signal => break,
631 _ = &mut timeout => {
632 continue;
635 }
636 res = frames.next() => {
637 match res {
638 Some(frame) => {
639 if let Some(permit) = &mut permit {
640 permit.decoding_finished(1);
644 };
645 handle_tcp_frame(&mut frame_handler, frame, &mut event_sink, received_from.clone(), Arc::clone(&active_parsing_task_nums)).await;
646 }
647 None => {
648 debug!("Connection closed.");
649 break
650 },
651 }
652 }
653 else => break,
654 }
655
656 drop(permit);
657 }
658}
659
660async fn handle_tcp_frame<T>(
661 frame_handler: &mut T,
662 frame: Bytes,
663 event_sink: &mut SourceSender,
664 received_from: Option<Bytes>,
665 active_parsing_task_nums: Arc<AtomicUsize>,
666) where
667 T: TcpFrameHandler + Send + Sync + Clone + 'static,
668{
669 if frame_handler.multithreaded() {
670 spawn_event_handling_tasks(
671 frame,
672 frame_handler.clone(),
673 event_sink.clone(),
674 received_from,
675 active_parsing_task_nums,
676 frame_handler.max_frame_handling_tasks(),
677 )
678 .await;
679 } else if let Some(event) = frame_handler.handle_event(received_from, frame)
680 && let Err(e) = event_sink.send_event(event).await
681 {
682 error!("Error sending event: {e:?}.");
683 }
684}
685
686pub fn build_framestream_unix_source(
692 frame_handler: impl UnixFrameHandler + Send + Sync + Clone + 'static,
693 shutdown: ShutdownSignal,
694 out: SourceSender,
695) -> crate::Result<Source> {
696 let path = frame_handler.socket_path();
697
698 match fs::metadata(&path) {
700 Ok(_) => {
701 info!(message = "Deleting file.", ?path);
703 fs::remove_file(&path)?;
704 }
705 Err(ref e) if e.kind() == std::io::ErrorKind::NotFound => {} Err(e) => {
707 error!("Unable to get socket information; error = {:?}.", e);
708 return Err(Box::new(e));
709 }
710 };
711
712 let listener = UnixListener::bind(&path)?;
713
714 if let Some(socket_receive_buffer_size) = frame_handler.socket_receive_buffer_size() {
716 _ = nix::sys::socket::setsockopt(
717 &listener,
718 nix::sys::socket::sockopt::RcvBuf,
719 &(socket_receive_buffer_size),
720 );
721 let rcv_buf_size =
722 nix::sys::socket::getsockopt(&listener, nix::sys::socket::sockopt::RcvBuf);
723 info!(
724 "Unix socket receive buffer size modified to {}.",
725 rcv_buf_size.unwrap()
726 );
727 }
728
729 if let Some(socket_send_buffer_size) = frame_handler.socket_send_buffer_size() {
731 _ = nix::sys::socket::setsockopt(
732 &listener,
733 nix::sys::socket::sockopt::SndBuf,
734 &(socket_send_buffer_size),
735 );
736 let snd_buf_size =
737 nix::sys::socket::getsockopt(&listener, nix::sys::socket::sockopt::SndBuf);
738 info!(
739 "Unix socket buffer send size modified to {}.",
740 snd_buf_size.unwrap()
741 );
742 }
743
744 if let Some(socket_permission) = frame_handler.socket_file_mode() {
746 if !(448..=511).contains(&socket_permission) {
747 return Err(format!(
748 "Invalid Socket permission {socket_permission:#o}. Must between 0o700 and 0o777."
749 )
750 .into());
751 }
752 match fs::set_permissions(&path, fs::Permissions::from_mode(socket_permission)) {
753 Ok(_) => {
754 info!("Socket permissions updated to {:#o}.", socket_permission);
755 }
756 Err(e) => {
757 error!(
758 "Failed to update listener socket permissions; error = {:?}.",
759 e
760 );
761 return Err(Box::new(e));
762 }
763 };
764 };
765
766 let fut = async move {
767 let active_parsing_task_nums = Arc::new(AtomicUsize::new(0));
768
769 info!(message = "Listening...", ?path, r#type = "unix");
770
771 let mut stream = UnixListenerStream::new(listener).take_until(shutdown.clone());
772 while let Some(socket) = stream.next().await {
773 let socket = match socket {
774 Err(e) => {
775 error!("Failed to accept socket; error = {:?}.", e);
776 continue;
777 }
778 Ok(s) => s,
779 };
780 let peer_addr = socket.peer_addr().ok();
781 let listen_path = path.clone();
782 let active_task_nums_ = Arc::clone(&active_parsing_task_nums);
783
784 let span = info_span!("connection");
785 let path = if let Some(addr) = peer_addr {
786 if let Some(path) = addr.as_pathname().map(|e| e.to_owned()) {
787 span.record("peer_path", field::debug(&path));
788 Some(path)
789 } else {
790 None
791 }
792 } else {
793 None
794 };
795 let received_from: Option<Bytes> =
796 path.map(|p| p.to_string_lossy().into_owned().into());
797
798 build_framestream_source(
799 frame_handler.clone(),
800 socket,
801 received_from,
802 out.clone(),
803 shutdown.clone(),
804 span,
805 active_task_nums_,
806 move |error| {
807 emit!(UnixSocketError {
808 error: &error,
809 path: &listen_path,
810 });
811 },
812 );
813 }
814
815 drop(stream);
817
818 if let Err(error) = fs::remove_file(&path) {
820 emit!(UnixSocketFileDeleteError { path: &path, error });
821 }
822
823 Ok(())
824 };
825
826 Ok(Box::pin(fut))
827}
828
829#[allow(clippy::too_many_arguments)]
830fn build_framestream_source<T: Send + 'static>(
831 frame_handler: impl FrameHandler + Send + Sync + Clone + 'static,
832 socket: impl AsyncRead + AsyncWrite + Send + 'static,
833 received_from: Option<Bytes>,
834 out: SourceSender,
835 shutdown: impl Future<Output = T> + Unpin + Send + 'static,
836 span: Span,
837 active_task_nums: Arc<AtomicUsize>,
838 error_mapper: impl FnMut(std::io::Error) + Send + 'static,
839) {
840 let content_type = frame_handler.content_type();
841 let mut event_sink = out.clone();
842 let (sock_sink, sock_stream) = Framed::new(
843 socket,
844 length_delimited::Builder::new()
845 .max_frame_length(frame_handler.max_frame_length())
846 .new_codec(),
847 )
848 .split();
849 let mut fs_reader = FrameStreamReader::new(Box::new(sock_sink), content_type);
850 let frame_handler_copy = frame_handler.clone();
851 let frames = sock_stream
852 .take_until(shutdown)
853 .map_err(error_mapper)
854 .filter_map(move |frame| {
855 future::ready(match frame {
856 Ok(f) => fs_reader.handle_frame(Bytes::from(f)),
857 Err(_) => None,
858 })
859 });
860 if !frame_handler.multithreaded() {
861 let mut events = frames.filter_map(move |f| {
862 future::ready(frame_handler_copy.handle_event(received_from.clone(), f))
863 });
864
865 let handler = async move {
866 if let Err(e) = event_sink.send_event_stream(&mut events).await {
867 error!("Error sending event: {:?}.", e);
868 }
869
870 info!("Finished sending.");
871 };
872 tokio::spawn(handler.instrument(span.or_current()));
873 } else {
874 let handler = async move {
875 frames
876 .for_each(move |f| {
877 let max_frame_handling_tasks = frame_handler_copy.max_frame_handling_tasks();
878 let f_handler = frame_handler_copy.clone();
879 let received_from_copy = received_from.clone();
880 let event_sink_copy = event_sink.clone();
881 let active_task_nums_copy = Arc::clone(&active_task_nums);
882
883 async move {
884 spawn_event_handling_tasks(
885 f,
886 f_handler,
887 event_sink_copy,
888 received_from_copy,
889 active_task_nums_copy,
890 max_frame_handling_tasks,
891 )
892 .await;
893 }
894 })
895 .await;
896 info!("Finished sending.");
897 };
898 tokio::spawn(handler.instrument(span.or_current()));
899 }
900}
901
902async fn spawn_event_handling_tasks(
903 event_data: Bytes,
904 event_handler: impl FrameHandler + Send + Sync + 'static,
905 mut event_sink: SourceSender,
906 received_from: Option<Bytes>,
907 active_task_nums: Arc<AtomicUsize>,
908 max_frame_handling_tasks: usize,
909) -> JoinHandle<()> {
910 wait_for_task_quota(&active_task_nums, max_frame_handling_tasks).await;
911
912 crate::spawn_in_current_span(async move {
913 future::ready({
914 if let Some(evt) = event_handler.handle_event(received_from, event_data)
915 && event_sink.send_event(evt).await.is_err()
916 {
917 error!("Encountered error while sending event.");
918 }
919 active_task_nums.fetch_sub(1, Ordering::AcqRel);
920 })
921 .await;
922 })
923}
924
925async fn wait_for_task_quota(active_task_nums: &Arc<AtomicUsize>, max_tasks: usize) {
926 while max_tasks > 0 && max_tasks < active_task_nums.load(Ordering::Acquire) {
927 tokio::time::sleep(Duration::from_millis(3)).await;
928 }
929 active_task_nums.fetch_add(1, Ordering::AcqRel);
930}
931
932#[cfg(test)]
933mod test {
934 use std::net::SocketAddr;
935 #[cfg(unix)]
936 use std::{
937 path::PathBuf,
938 sync::{
939 Arc,
940 atomic::{AtomicUsize, Ordering},
941 },
942 thread,
943 };
944
945 use bytes::{Bytes, BytesMut, buf::Buf};
946 use futures::{
947 future,
948 sink::{Sink, SinkExt},
949 stream::{self, StreamExt},
950 };
951 use futures_util::Stream;
952 use ipnet::IpNet;
953 use tokio::{
954 self,
955 net::{TcpStream, UnixStream},
956 task::JoinHandle,
957 time::{Duration, Instant},
958 };
959 use tokio_util::codec::{Framed, length_delimited};
960 use vector_lib::{
961 config::{LegacyKey, LogNamespace},
962 lookup::{OwnedValuePath, owned_value_path, path},
963 tcp::TcpKeepaliveConfig,
964 tls::{CertificateMetadata, MaybeTls, MaybeTlsSettings},
965 };
966
967 use super::{
968 ControlField, ControlHeader, FrameHandler, TcpFrameHandler, UnixFrameHandler,
969 build_framestream_tcp_source, build_framestream_unix_source, spawn_event_handling_tasks,
970 };
971 use crate::{
972 SourceSender,
973 config::{ComponentKey, log_schema},
974 event::{Event, LogEvent},
975 shutdown::SourceShutdownCoordinator,
976 sources::util::net::SocketListenAddr,
977 test_util::{addr::next_addr, collect_n, collect_n_stream},
978 };
979
980 #[derive(Clone)]
981 struct MockFrameHandler<F: Send + Sync + Clone + FnOnce() + 'static> {
982 content_type: String,
983 max_frame_length: usize,
984 multithreaded: bool,
985 max_frame_handling_tasks: usize,
986 extra_task_handling_routine: F,
987 host_key: Option<OwnedValuePath>,
988 timestamp_key: Option<OwnedValuePath>,
989 source_type_key: Option<OwnedValuePath>,
990 log_namespace: LogNamespace,
991 }
992
993 #[derive(Clone)]
994 struct MockUnixFrameHandler<F: Send + Sync + Clone + FnOnce() + 'static> {
995 frame_handler: MockFrameHandler<F>,
996 socket_path: PathBuf,
997 socket_file_mode: Option<u32>,
998 socket_receive_buffer_size: Option<usize>,
999 socket_send_buffer_size: Option<usize>,
1000 }
1001
1002 #[derive(Clone)]
1003 struct MockTcpFrameHandler<F: Send + Sync + Clone + FnOnce() + 'static> {
1004 frame_handler: MockFrameHandler<F>,
1005 address: SocketListenAddr,
1006 keepalive: Option<TcpKeepaliveConfig>,
1007 shutdown_timeout_secs: Duration,
1008 tls: MaybeTlsSettings,
1009 tls_client_metadata_key: Option<OwnedValuePath>,
1010 receive_buffer_bytes: Option<usize>,
1011 max_connection_duration_secs: Option<u64>,
1012 max_connections: Option<u32>,
1013 permit_origin: Option<Vec<IpNet>>,
1014 }
1015
1016 impl<F: Send + Sync + Clone + FnOnce() + 'static> MockTcpFrameHandler<F> {
1017 pub fn new(
1018 addr: SocketAddr,
1019 content_type: String,
1020 multithreaded: bool,
1021 extra_routine: F,
1022 permit_origin: Option<Vec<IpNet>>,
1023 ) -> Self {
1024 Self {
1025 frame_handler: MockFrameHandler::new(content_type, multithreaded, extra_routine),
1026 address: addr.into(),
1027 keepalive: None,
1028 shutdown_timeout_secs: Duration::from_secs(30),
1029 tls: MaybeTls::Raw(()),
1030 tls_client_metadata_key: None,
1031 receive_buffer_bytes: None,
1032 max_connection_duration_secs: None,
1033 max_connections: None,
1034 permit_origin,
1035 }
1036 }
1037 }
1038
1039 impl<F: Send + Sync + Clone + FnOnce() + 'static> MockUnixFrameHandler<F> {
1040 pub fn new(content_type: String, multithreaded: bool, extra_routine: F) -> Self {
1041 Self {
1042 frame_handler: MockFrameHandler::new(content_type, multithreaded, extra_routine),
1043 socket_path: tempfile::tempdir().unwrap().keep().join("unix_test"),
1044 socket_file_mode: None,
1045 socket_receive_buffer_size: None,
1046 socket_send_buffer_size: None,
1047 }
1048 }
1049 }
1050
1051 impl<F: Send + Sync + Clone + FnOnce() + 'static> MockFrameHandler<F> {
1052 pub fn new(content_type: String, multithreaded: bool, extra_routine: F) -> Self {
1053 Self {
1054 content_type,
1055 max_frame_length: bytesize::kib(100u64) as usize,
1056 multithreaded,
1057 max_frame_handling_tasks: 0,
1058 extra_task_handling_routine: extra_routine,
1059 host_key: Some(owned_value_path!("test_framestream")),
1060 timestamp_key: Some(owned_value_path!("my_timestamp")),
1061 source_type_key: Some(owned_value_path!("source_type")),
1062 log_namespace: LogNamespace::Legacy,
1063 }
1064 }
1065 }
1066
1067 impl<F: Send + Sync + Clone + FnOnce() + 'static> FrameHandler for MockFrameHandler<F> {
1068 fn content_type(&self) -> String {
1069 self.content_type.clone()
1070 }
1071 fn max_frame_length(&self) -> usize {
1072 self.max_frame_length
1073 }
1074
1075 fn handle_event(&self, received_from: Option<Bytes>, frame: Bytes) -> Option<Event> {
1076 let mut log_event = LogEvent::from(frame);
1077
1078 log_event.insert(
1079 log_schema().source_type_key_target_path().unwrap(),
1080 "framestream",
1081 );
1082 if let Some(host) = received_from {
1083 self.log_namespace.insert_source_metadata(
1084 "framestream",
1085 &mut log_event,
1086 self.host_key.as_ref().map(LegacyKey::Overwrite),
1087 path!("host"),
1088 host,
1089 )
1090 }
1091
1092 (self.extra_task_handling_routine.clone())();
1093
1094 Some(log_event.into())
1095 }
1096
1097 fn multithreaded(&self) -> bool {
1098 self.multithreaded
1099 }
1100 fn max_frame_handling_tasks(&self) -> usize {
1101 self.max_frame_handling_tasks
1102 }
1103
1104 fn host_key(&self) -> &Option<OwnedValuePath> {
1105 &self.host_key
1106 }
1107
1108 fn timestamp_key(&self) -> Option<&OwnedValuePath> {
1109 self.timestamp_key.as_ref()
1110 }
1111
1112 fn source_type_key(&self) -> Option<&OwnedValuePath> {
1113 self.source_type_key.as_ref()
1114 }
1115 }
1116
1117 impl<F: Send + Sync + Clone + FnOnce() + 'static> FrameHandler for MockUnixFrameHandler<F> {
1118 fn content_type(&self) -> String {
1119 self.frame_handler.content_type()
1120 }
1121
1122 fn max_frame_length(&self) -> usize {
1123 self.frame_handler.max_frame_length()
1124 }
1125
1126 fn handle_event(&self, received_from: Option<Bytes>, frame: Bytes) -> Option<Event> {
1127 self.frame_handler.handle_event(received_from, frame)
1128 }
1129
1130 fn multithreaded(&self) -> bool {
1131 self.frame_handler.multithreaded()
1132 }
1133
1134 fn max_frame_handling_tasks(&self) -> usize {
1135 self.frame_handler.max_frame_handling_tasks()
1136 }
1137
1138 fn host_key(&self) -> &Option<OwnedValuePath> {
1139 self.frame_handler.host_key()
1140 }
1141
1142 fn timestamp_key(&self) -> Option<&OwnedValuePath> {
1143 self.frame_handler.timestamp_key()
1144 }
1145
1146 fn source_type_key(&self) -> Option<&OwnedValuePath> {
1147 self.frame_handler.source_type_key()
1148 }
1149 }
1150
1151 impl<F: Send + Sync + Clone + FnOnce() + 'static> UnixFrameHandler for MockUnixFrameHandler<F> {
1152 fn socket_path(&self) -> PathBuf {
1153 self.socket_path.clone()
1154 }
1155
1156 fn socket_file_mode(&self) -> Option<u32> {
1157 self.socket_file_mode
1158 }
1159
1160 fn socket_receive_buffer_size(&self) -> Option<usize> {
1161 self.socket_receive_buffer_size
1162 }
1163
1164 fn socket_send_buffer_size(&self) -> Option<usize> {
1165 self.socket_send_buffer_size
1166 }
1167 }
1168
1169 impl<F: Send + Sync + Clone + FnOnce() + 'static> FrameHandler for MockTcpFrameHandler<F> {
1170 fn content_type(&self) -> String {
1171 self.frame_handler.content_type()
1172 }
1173
1174 fn max_frame_length(&self) -> usize {
1175 self.frame_handler.max_frame_length()
1176 }
1177
1178 fn handle_event(&self, received_from: Option<Bytes>, frame: Bytes) -> Option<Event> {
1179 self.frame_handler.handle_event(received_from, frame)
1180 }
1181
1182 fn multithreaded(&self) -> bool {
1183 self.frame_handler.multithreaded()
1184 }
1185
1186 fn max_frame_handling_tasks(&self) -> usize {
1187 self.frame_handler.max_frame_handling_tasks()
1188 }
1189
1190 fn host_key(&self) -> &Option<OwnedValuePath> {
1191 self.frame_handler.host_key()
1192 }
1193
1194 fn timestamp_key(&self) -> Option<&OwnedValuePath> {
1195 self.frame_handler.timestamp_key()
1196 }
1197
1198 fn source_type_key(&self) -> Option<&OwnedValuePath> {
1199 self.frame_handler.source_type_key()
1200 }
1201 }
1202
1203 impl<F: Send + Sync + Clone + FnOnce() + 'static> TcpFrameHandler for MockTcpFrameHandler<F> {
1204 fn address(&self) -> SocketListenAddr {
1205 self.address
1206 }
1207
1208 fn keepalive(&self) -> Option<TcpKeepaliveConfig> {
1209 self.keepalive
1210 }
1211
1212 fn shutdown_timeout_secs(&self) -> Duration {
1213 self.shutdown_timeout_secs
1214 }
1215
1216 fn tls(&self) -> MaybeTlsSettings {
1217 self.tls.clone()
1218 }
1219
1220 fn tls_client_metadata_key(&self) -> Option<OwnedValuePath> {
1221 self.tls_client_metadata_key.clone()
1222 }
1223
1224 fn receive_buffer_bytes(&self) -> Option<usize> {
1225 self.receive_buffer_bytes
1226 }
1227
1228 fn max_connection_duration_secs(&self) -> Option<u64> {
1229 self.max_connection_duration_secs
1230 }
1231
1232 fn max_connections(&self) -> Option<u32> {
1233 self.max_connections
1234 }
1235
1236 fn insert_tls_client_metadata(&mut self, _: Option<CertificateMetadata>) {}
1237
1238 fn allowed_origins(&self) -> Option<&[IpNet]> {
1239 self.permit_origin.as_deref()
1240 }
1241 }
1242
1243 fn init_framestream_tcp(
1244 source_id: &str,
1245 addr: &SocketAddr,
1246 frame_handler: impl TcpFrameHandler + Send + Sync + Clone + 'static,
1247 pipeline: SourceSender,
1248 ) -> (JoinHandle<Result<(), ()>>, SourceShutdownCoordinator) {
1249 let source_id = ComponentKey::from(source_id);
1250 let mut shutdown = SourceShutdownCoordinator::default();
1251 let (shutdown_signal, _) = shutdown.register_source(&source_id, false);
1252 let server = build_framestream_tcp_source(frame_handler, shutdown_signal, pipeline)
1253 .expect("Failed to build framestream tcp source.");
1254
1255 let join_handle = tokio::spawn(server);
1256
1257 while std::net::TcpStream::connect(addr).is_err() {
1258 thread::sleep(Duration::from_millis(2));
1259 }
1260
1261 (join_handle, shutdown)
1262 }
1263
1264 fn init_framestream_unix(
1265 source_id: &str,
1266 frame_handler: impl UnixFrameHandler + Send + Sync + Clone + 'static,
1267 pipeline: SourceSender,
1268 ) -> (
1269 PathBuf,
1270 JoinHandle<Result<(), ()>>,
1271 SourceShutdownCoordinator,
1272 ) {
1273 let source_id = ComponentKey::from(source_id);
1274 let socket_path = frame_handler.socket_path();
1275 let mut shutdown = SourceShutdownCoordinator::default();
1276 let (shutdown_signal, _) = shutdown.register_source(&source_id, false);
1277 let server = build_framestream_unix_source(frame_handler, shutdown_signal, pipeline)
1278 .expect("Failed to build framestream unix source.");
1279
1280 let join_handle = tokio::spawn(server);
1281
1282 while std::os::unix::net::UnixStream::connect(&socket_path).is_err() {
1284 thread::sleep(Duration::from_millis(2));
1285 }
1286
1287 (socket_path, join_handle, shutdown)
1288 }
1289
1290 async fn make_tcp_stream(
1291 addr: SocketAddr,
1292 ) -> Framed<TcpStream, length_delimited::LengthDelimitedCodec> {
1293 let socket = TcpStream::connect(&addr).await.unwrap();
1294 Framed::new(socket, length_delimited::Builder::new().new_codec())
1295 }
1296
1297 async fn make_unix_stream(
1298 path: PathBuf,
1299 ) -> Framed<UnixStream, length_delimited::LengthDelimitedCodec> {
1300 let socket = UnixStream::connect(&path).await.unwrap();
1301 Framed::new(socket, length_delimited::Builder::new().new_codec())
1302 }
1303
1304 async fn send_data_frames<S: Sink<Bytes, Error = std::io::Error> + Unpin>(
1305 sock_sink: &mut S,
1306 frames: Vec<Result<Bytes, std::io::Error>>,
1307 ) {
1308 let mut stream = stream::iter(frames);
1309 _ = sock_sink.send_all(&mut stream).await;
1311 }
1312
1313 async fn send_control_frame<S: Sink<Bytes, Error = std::io::Error> + Unpin>(
1314 sock_sink: &mut S,
1315 frame: Bytes,
1316 ) {
1317 send_data_frames(sock_sink, vec![Ok(Bytes::new()), Ok(frame)]).await; }
1319
1320 fn create_control_frame(header: ControlHeader) -> Bytes {
1321 Bytes::from(header.to_u32().to_be_bytes().to_vec())
1322 }
1323
1324 fn create_control_frame_with_content(
1325 header: ControlHeader,
1326 content_types: Vec<Bytes>,
1327 ) -> Bytes {
1328 let mut frame = BytesMut::from(&header.to_u32().to_be_bytes()[..]);
1329 for content_type in content_types {
1330 frame.extend(ControlField::ContentType.to_u32().to_be_bytes());
1331 frame.extend((content_type.len() as u32).to_be_bytes());
1332 frame.extend(content_type.clone());
1333 }
1334 Bytes::from(frame)
1335 }
1336
1337 fn assert_accept_frame(frame: &mut BytesMut, expected_content_type: Bytes) {
1338 assert_eq!(&frame[..4], &ControlHeader::Accept.to_u32().to_be_bytes(),);
1341 frame.advance(4);
1342 assert_eq!(
1344 &frame[..4],
1345 &ControlField::ContentType.to_u32().to_be_bytes(),
1346 );
1347 frame.advance(4);
1348 assert_eq!(
1350 &frame[..4],
1351 &(expected_content_type.len() as u32).to_be_bytes(),
1352 );
1353 frame.advance(4);
1354 assert_eq!(&frame[..], &expected_content_type[..]);
1356 }
1357
1358 fn create_frame_handler(multithreaded: bool) -> impl UnixFrameHandler + Send + Sync + Clone {
1359 MockUnixFrameHandler::new("test_content".to_string(), multithreaded, move || {})
1360 }
1361
1362 fn create_tcp_frame_handler(
1363 addr: SocketAddr,
1364 multithreaded: bool,
1365 permit_origin: Option<Vec<IpNet>>,
1366 ) -> impl TcpFrameHandler + Send + Sync + Clone {
1367 MockTcpFrameHandler::new(
1368 addr,
1369 "test_content".to_string(),
1370 multithreaded,
1371 move || {},
1372 permit_origin,
1373 )
1374 }
1375
1376 async fn signal_shutdown(source_name: &str, shutdown: &mut SourceShutdownCoordinator) {
1377 let deadline = Instant::now() + Duration::from_secs(10);
1379 let id = ComponentKey::from(source_name);
1380 let shutdown_complete = shutdown.shutdown_source(&id, deadline);
1381 let shutdown_success = shutdown_complete.await;
1382 assert!(shutdown_success);
1383 }
1384
1385 async fn test_normal_framestream<
1386 T: Sink<Bytes, Error = std::io::Error> + Unpin,
1387 U: Stream<Item = Result<BytesMut, std::io::Error>> + Unpin,
1388 V: Stream<Item = Event> + Unpin,
1389 >(
1390 source_name: &str,
1391 mut sock_sink: T,
1392 mut sock_stream: U,
1393 rx: V,
1394 mut shutdown: SourceShutdownCoordinator,
1395 source_handle: JoinHandle<Result<(), ()>>,
1396 ) {
1397 let content_type = Bytes::from(&b"test_content"[..]);
1399 let ready_msg =
1400 create_control_frame_with_content(ControlHeader::Ready, vec![content_type.clone()]);
1401 send_control_frame(&mut sock_sink, ready_msg).await;
1402
1403 let mut frame_vec = collect_n_stream(&mut sock_stream, 2).await;
1405 assert_eq!(frame_vec[0].as_ref().unwrap().len(), 0);
1407 assert_accept_frame(frame_vec[1].as_mut().unwrap(), content_type);
1408
1409 send_control_frame(&mut sock_sink, create_control_frame(ControlHeader::Start)).await;
1411
1412 send_data_frames(
1414 &mut sock_sink,
1415 vec![Ok(Bytes::from("hello")), Ok(Bytes::from("world"))],
1416 )
1417 .await;
1418 let events = collect_n(rx, 2).await;
1419
1420 send_control_frame(&mut sock_sink, create_control_frame(ControlHeader::Stop)).await;
1422
1423 let message_key = log_schema().message_key().unwrap().to_string();
1424 assert!(
1425 events
1426 .iter()
1427 .any(|e| e.as_log()[&message_key] == "hello".into())
1428 );
1429 assert!(
1430 events
1431 .iter()
1432 .any(|e| e.as_log()[&message_key] == "world".into())
1433 );
1434
1435 drop(sock_stream); signal_shutdown(source_name, &mut shutdown).await;
1439 _ = source_handle.await.unwrap();
1440 }
1441
1442 async fn test_multiple_content_types<
1443 T: Sink<Bytes, Error = std::io::Error> + Unpin,
1444 U: Stream<Item = Result<BytesMut, std::io::Error>> + Unpin,
1445 >(
1446 source_name: &str,
1447 mut sock_sink: T,
1448 mut sock_stream: U,
1449 mut shutdown: SourceShutdownCoordinator,
1450 source_handle: JoinHandle<Result<(), ()>>,
1451 ) {
1452 let content_type = Bytes::from(&b"test_content"[..]);
1454 let ready_msg = create_control_frame_with_content(
1455 ControlHeader::Ready,
1456 vec![Bytes::from(&b"test_content2"[..]), content_type.clone()],
1457 ); send_control_frame(&mut sock_sink, ready_msg).await;
1459
1460 let mut frame_vec = collect_n_stream(&mut sock_stream, 2).await;
1462
1463 assert_eq!(frame_vec[0].as_ref().unwrap().len(), 0);
1465 assert_accept_frame(frame_vec[1].as_mut().unwrap(), content_type);
1466
1467 drop(sock_stream); signal_shutdown(source_name, &mut shutdown).await;
1471 _ = source_handle.await.unwrap();
1472 }
1473
1474 #[tokio::test(flavor = "multi_thread")]
1475 #[should_panic]
1476 async fn blocked_framestream_tcp() {
1477 let source_name = "test_source";
1478 let (tx, rx) = SourceSender::new_test();
1479 let (_guard, addr) = next_addr();
1480 let (source_handle, shutdown) = init_framestream_tcp(
1481 source_name,
1482 &addr,
1483 create_tcp_frame_handler(addr, false, Some(vec![])),
1484 tx,
1485 );
1486 let (sock_sink, sock_stream) = make_tcp_stream(addr).await.split();
1487
1488 test_normal_framestream(
1489 source_name,
1490 sock_sink,
1491 sock_stream,
1492 rx,
1493 shutdown,
1494 source_handle,
1495 )
1496 .await;
1497 }
1498
1499 #[tokio::test(flavor = "multi_thread")]
1500 async fn normal_framestream_singlethreaded_tcp() {
1501 let source_name = "test_source";
1502 let (tx, rx) = SourceSender::new_test();
1503 let (_guard, addr) = next_addr();
1504 let (source_handle, shutdown) = init_framestream_tcp(
1505 source_name,
1506 &addr,
1507 create_tcp_frame_handler(addr, false, None),
1508 tx,
1509 );
1510 let (sock_sink, sock_stream) = make_tcp_stream(addr).await.split();
1511
1512 test_normal_framestream(
1513 source_name,
1514 sock_sink,
1515 sock_stream,
1516 rx,
1517 shutdown,
1518 source_handle,
1519 )
1520 .await;
1521 }
1522
1523 #[tokio::test(flavor = "multi_thread")]
1524 async fn normal_framestream_singlethreaded_unix() {
1525 let source_name = "test_source";
1526 let (tx, rx) = SourceSender::new_test();
1527 let (path, source_handle, shutdown) =
1528 init_framestream_unix(source_name, create_frame_handler(false), tx);
1529 let (sock_sink, sock_stream) = make_unix_stream(path).await.split();
1530
1531 test_normal_framestream(
1532 source_name,
1533 sock_sink,
1534 sock_stream,
1535 rx,
1536 shutdown,
1537 source_handle,
1538 )
1539 .await;
1540 }
1541
1542 #[tokio::test(flavor = "multi_thread")]
1543 async fn normal_framestream_multithreaded_tcp() {
1544 let source_name = "test_source";
1545 let (tx, rx) = SourceSender::new_test();
1546 let (_guard, addr) = next_addr();
1547 let (source_handle, shutdown) = init_framestream_tcp(
1548 source_name,
1549 &addr,
1550 create_tcp_frame_handler(addr, true, None),
1551 tx,
1552 );
1553 let (sock_sink, sock_stream) = make_tcp_stream(addr).await.split();
1554
1555 test_normal_framestream(
1556 source_name,
1557 sock_sink,
1558 sock_stream,
1559 rx,
1560 shutdown,
1561 source_handle,
1562 )
1563 .await;
1564 }
1565
1566 #[tokio::test(flavor = "multi_thread")]
1567 async fn normal_framestream_multithreaded_unix() {
1568 let source_name = "test_source";
1569 let (tx, rx) = SourceSender::new_test();
1570 let (path, source_handle, shutdown) =
1571 init_framestream_unix(source_name, create_frame_handler(true), tx);
1572 let (sock_sink, sock_stream) = make_unix_stream(path).await.split();
1573
1574 test_normal_framestream(
1575 source_name,
1576 sock_sink,
1577 sock_stream,
1578 rx,
1579 shutdown,
1580 source_handle,
1581 )
1582 .await;
1583 }
1584
1585 #[tokio::test(flavor = "multi_thread")]
1586 async fn multiple_content_types_tcp() {
1587 let source_name = "test_source";
1588 let (tx, _) = SourceSender::new_test();
1589 let (_guard, addr) = next_addr();
1590 let (source_handle, shutdown) = init_framestream_tcp(
1591 source_name,
1592 &addr,
1593 create_tcp_frame_handler(addr, false, None),
1594 tx,
1595 );
1596 let (sock_sink, sock_stream) = make_tcp_stream(addr).await.split();
1597
1598 test_multiple_content_types(source_name, sock_sink, sock_stream, shutdown, source_handle)
1599 .await;
1600 }
1601
1602 #[tokio::test(flavor = "multi_thread")]
1603 async fn multiple_content_types_unix() {
1604 let source_name = "test_source";
1605 let (tx, _) = SourceSender::new_test();
1606 let (path, source_handle, shutdown) =
1607 init_framestream_unix(source_name, create_frame_handler(false), tx);
1608 let (sock_sink, sock_stream) = make_unix_stream(path).await.split();
1609
1610 test_multiple_content_types(source_name, sock_sink, sock_stream, shutdown, source_handle)
1611 .await;
1612 }
1613
1614 #[tokio::test(flavor = "multi_thread")]
1615 async fn wrong_content_type() {
1616 let source_name = "test_source";
1617 let (tx, _) = SourceSender::new_test();
1618 let (path, source_handle, mut shutdown) =
1619 init_framestream_unix(source_name, create_frame_handler(false), tx);
1620 let (mut sock_sink, mut sock_stream) = make_unix_stream(path).await.split();
1621
1622 let ready_msg = create_control_frame_with_content(
1624 ControlHeader::Ready,
1625 vec![Bytes::from(&b"test_content2"[..])],
1626 ); send_control_frame(&mut sock_sink, ready_msg).await;
1628
1629 let content_type = Bytes::from(&b"test_content"[..]);
1631 let ready_msg =
1632 create_control_frame_with_content(ControlHeader::Ready, vec![content_type.clone()]);
1633 send_control_frame(&mut sock_sink, ready_msg).await;
1634
1635 let mut frame_vec = collect_n_stream(&mut sock_stream, 2).await;
1637
1638 assert_eq!(frame_vec[0].as_ref().unwrap().len(), 0);
1640 assert_accept_frame(frame_vec[1].as_mut().unwrap(), content_type);
1641
1642 drop(sock_stream); signal_shutdown(source_name, &mut shutdown).await;
1646 _ = source_handle.await.unwrap();
1647 }
1648
1649 #[tokio::test(flavor = "multi_thread")]
1650 async fn data_too_soon() {
1651 let source_name = "test_source";
1652 let (tx, rx) = SourceSender::new_test();
1653 let (path, source_handle, mut shutdown) =
1654 init_framestream_unix(source_name, create_frame_handler(false), tx);
1655 let (mut sock_sink, mut sock_stream) = make_unix_stream(path).await.split();
1656
1657 send_data_frames(
1659 &mut sock_sink,
1660 vec![Ok(Bytes::from("bad")), Ok(Bytes::from("data"))],
1661 )
1662 .await;
1663
1664 let content_type = Bytes::from(&b"test_content"[..]);
1666 let ready_msg =
1667 create_control_frame_with_content(ControlHeader::Ready, vec![content_type.clone()]);
1668 send_control_frame(&mut sock_sink, ready_msg).await;
1669
1670 let mut frame_vec = collect_n_stream(&mut sock_stream, 2).await;
1672
1673 assert_eq!(frame_vec[0].as_ref().unwrap().len(), 0);
1675 assert_accept_frame(frame_vec[1].as_mut().unwrap(), content_type);
1676
1677 send_control_frame(&mut sock_sink, create_control_frame(ControlHeader::Start)).await;
1679
1680 send_data_frames(
1682 &mut sock_sink,
1683 vec![Ok(Bytes::from("hello")), Ok(Bytes::from("world"))],
1684 )
1685 .await;
1686 let events = collect_n(rx, 2).await;
1687
1688 send_control_frame(&mut sock_sink, create_control_frame(ControlHeader::Stop)).await;
1690
1691 assert_eq!(
1692 events[0].as_log()[log_schema().message_key().unwrap().to_string()],
1693 "hello".into(),
1694 );
1695 assert_eq!(
1696 events[1].as_log()[log_schema().message_key().unwrap().to_string()],
1697 "world".into(),
1698 );
1699
1700 drop(sock_stream); signal_shutdown(source_name, &mut shutdown).await;
1704 _ = source_handle.await.unwrap();
1705 }
1706
1707 #[tokio::test(flavor = "multi_thread")]
1708 async fn unidirectional_framestream() {
1709 let source_name = "test_source";
1710 let (tx, rx) = SourceSender::new_test();
1711 let (path, source_handle, mut shutdown) =
1712 init_framestream_unix(source_name, create_frame_handler(false), tx);
1713 let (mut sock_sink, _) = make_unix_stream(path).await.split();
1714
1715 let content_type = Bytes::from(&b"test_content"[..]);
1717 let start_msg = create_control_frame_with_content(ControlHeader::Start, vec![content_type]);
1718 send_control_frame(&mut sock_sink, start_msg).await;
1719
1720 send_data_frames(
1722 &mut sock_sink,
1723 vec![Ok(Bytes::from("hello")), Ok(Bytes::from("world"))],
1724 )
1725 .await;
1726 let events = collect_n(rx, 2).await;
1727
1728 send_control_frame(&mut sock_sink, create_control_frame(ControlHeader::Stop)).await;
1730
1731 assert_eq!(
1732 events[0].as_log()[log_schema().message_key().unwrap().to_string()],
1733 "hello".into(),
1734 );
1735 assert_eq!(
1736 events[1].as_log()[log_schema().message_key().unwrap().to_string()],
1737 "world".into(),
1738 );
1739
1740 signal_shutdown(source_name, &mut shutdown).await;
1742 _ = source_handle.await.unwrap();
1743 }
1744
1745 #[tokio::test(flavor = "multi_thread")]
1746 async fn test_spawn_event_handling_tasks() {
1747 let (out, rx) = SourceSender::new_test();
1748
1749 let max_frame_handling_tasks = 20;
1750 let active_task_nums = Arc::new(AtomicUsize::new(0));
1751 let active_task_nums_copy = Arc::clone(&active_task_nums);
1752 let max_task_nums_reached = Arc::new(AtomicUsize::new(0));
1753 let max_task_nums_reached_copy = Arc::clone(&max_task_nums_reached);
1754
1755 let mut join_handles = vec![];
1756 let active_task_nums_copy_2 = Arc::clone(&active_task_nums_copy);
1757 let extra_routine = move || {
1758 thread::sleep(Duration::from_millis(10));
1759 max_task_nums_reached_copy.fetch_max(
1760 active_task_nums_copy_2.load(Ordering::Acquire),
1761 Ordering::AcqRel,
1762 );
1763 };
1764
1765 let total_events = max_frame_handling_tasks * 10;
1766
1767 join_handles.push(tokio::spawn(async move {
1768 future::ready({
1769 let events = collect_n(rx, total_events).await;
1770 assert_eq!(total_events, events.len(), "Missed events");
1771 })
1772 .await;
1773 }));
1774
1775 for i in 0..total_events {
1776 join_handles.push(
1777 spawn_event_handling_tasks(
1778 Bytes::from(format!("event_{i}")),
1779 MockFrameHandler::new("test_content".to_string(), true, extra_routine.clone()),
1780 out.clone(),
1781 None,
1782 Arc::clone(&active_task_nums_copy),
1783 max_frame_handling_tasks,
1784 )
1785 .await,
1786 );
1787 }
1788
1789 future::join_all(join_handles).await;
1790
1791 let final_task_nums = active_task_nums.load(Ordering::Acquire);
1792 assert_eq!(
1793 0, final_task_nums,
1794 "There should be NO left-over tasks at the end"
1795 );
1796
1797 let max_task_nums_reached_value = max_task_nums_reached.load(Ordering::Acquire);
1798 assert!(
1799 max_task_nums_reached_value > 1,
1800 "MultiThreaded mode does NOT work"
1801 );
1802 assert!(
1803 (max_task_nums_reached_value - max_frame_handling_tasks) < 2,
1804 "Max number of tasks at any given time should NOT Exceed max_frame_handling_tasks too much"
1805 );
1806 }
1807}