Skip to main content

vector/sources/util/
framestream.rs

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
64/// If a connection does not receive any data during this short timeout,
65/// it should release its permit (and try to obtain a new one) allowing other connections to read.
66/// It is very short because any incoming data will avoid this timeout,
67/// so it mainly prevents holding permits without consuming any data
68const 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            //first control frame should be READY (if bidirectional -- if unidirectional first will be START)
88            control_state: ControlState::Initial,
89            is_bidirectional: true, //assume
90        }
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            //frame length of zero means the next frame is a control frame
179            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            //data frame
187            if self.state.control_state == ControlState::ReadingData {
188                Some(frame) //return data frame
189            } 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        //enforce maximum control frame size
201        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 current state to received header
208        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; //waiting for a START control frame
219                    }
220                    ControlHeader::Start => {
221                        //check for content type
222                        _ = self.process_fields(header, &mut frame)?;
223                        //if didn't error, then we are ok to change state
224                        self.state.control_state = ControlState::ReadingData;
225                        self.state.is_bidirectional = false; //if first message was START then we are unidirectional (no responses)
226                    }
227                    _ => error!("Got wrong control frame, expected READY."),
228                }
229            }
230            ControlState::GotReady => {
231                match header {
232                    ControlHeader::Start => {
233                        //check for content type
234                        _ = self.process_fields(header, &mut frame)?;
235                        //if didn't error, then we are ok to change state
236                        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                        //check there aren't any fields
245                        _ = self.process_fields(header, &mut frame)?;
246                        if self.state.is_bidirectional {
247                            //send FINISH frame -- but only if we are bidirectional
248                            self.send_control_frame(Self::make_frame(ControlHeader::Finish, None));
249                        }
250                        self.state.control_state = ControlState::Stopped; //stream is now done
251                    }
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                //should provide 1+ content types
268                //should match expected content type
269                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                //can take one or zero content types
275                if frame.is_empty() {
276                    Ok(None)
277                } else {
278                    //should match expected content type
279                    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                //check that there are no fields
286                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            //4 bytes of ControlField
309            let field_val = advance_u32(frame)?;
310            let field_type = ControlField::from_u32(field_val)?;
311            match field_type {
312                ControlField::ContentType => {
313                    //4 bytes giving length of content type
314                    let field_len = advance_u32(frame)? as usize;
315
316                    //enforce limit on content type string
317                    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()); //field type: ContentType
355            frame.extend((s.len() as u32).to_be_bytes()); //length of type
356            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""[..]); //send empty frame to say we are control frame
363        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
402/**
403 * Based off of the build_framestream_unix_source function.
404 * Functions similarly, just uses TCP socket instead of unix socket
405 **/
406pub 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, // tls_reloader: not wired for this source
423            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                // This connection is currently holding a permit, but has not received data for some time. Release
633                // the permit to let another connection try
634                continue;
635            }
636            res = frames.next() => {
637                match res {
638                    Some(frame) => {
639                        if let Some(permit) = &mut permit {
640                            // Note that this is intentionally not the "number of events in a single request", but rather
641                            // the "number of events currently available". This may contain events from multiple events,
642                            // but it should always contain all events from each request.
643                            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
686/**
687 * Based off of the build_unix_source function.
688 * Functions similarly, but uses the FrameStreamReader to deal with
689 * framestream control packets, and responds appropriately.
690 **/
691pub 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    //check if the path already exists (and try to delete it)
699    match fs::metadata(&path) {
700        Ok(_) => {
701            //exists, so try to delete it
702            info!(message = "Deleting file.", ?path);
703            fs::remove_file(&path)?;
704        }
705        Err(ref e) if e.kind() == std::io::ErrorKind::NotFound => {} //doesn't exist, do nothing
706        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    // system's 'net.core.rmem_max' might have to be changed if socket receive buffer is not updated properly
715    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    // system's 'net.core.wmem_max' might have to be changed if socket send buffer is not updated properly
730    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    // the permissions to unix socket are restricted from 0o700 to 0o777, which are 448 and 511 in decimal
745    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        // Cleanup
816        drop(stream);
817
818        // Delete socket file
819        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        // Wait for server to accept traffic
1283        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        //send and send_all consume the sink
1310        _ = 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; //send empty frame to say we are control frame
1318    }
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        //frame should start with 4 bytes saying ACCEPT
1339
1340        assert_eq!(&frame[..4], &ControlHeader::Accept.to_u32().to_be_bytes(),);
1341        frame.advance(4);
1342        //next should be content type field
1343        assert_eq!(
1344            &frame[..4],
1345            &ControlField::ContentType.to_u32().to_be_bytes(),
1346        );
1347        frame.advance(4);
1348        //next should be length of content_type
1349        assert_eq!(
1350            &frame[..4],
1351            &(expected_content_type.len() as u32).to_be_bytes(),
1352        );
1353        frame.advance(4);
1354        //rest should be content type
1355        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        // Now signal to the Source to shut down.
1378        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        //1 - send READY frame (with content_type)
1398        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        //2 - wait for ACCEPT frame
1404        let mut frame_vec = collect_n_stream(&mut sock_stream, 2).await;
1405        //take second element, because first will be empty (signifying control frame)
1406        assert_eq!(frame_vec[0].as_ref().unwrap().len(), 0);
1407        assert_accept_frame(frame_vec[1].as_mut().unwrap(), content_type);
1408
1409        //3 - send START frame
1410        send_control_frame(&mut sock_sink, create_control_frame(ControlHeader::Start)).await;
1411
1412        //4 - send data
1413        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        //5 - send STOP frame
1421        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); //explicitly drop the stream so we don't get warnings about not using it
1436
1437        // Ensure source actually shut down successfully.
1438        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        //1 - send READY frame (with content_type)
1453        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        ); //can have multiple content types
1458        send_control_frame(&mut sock_sink, ready_msg).await;
1459
1460        //2 - wait for ACCEPT frame
1461        let mut frame_vec = collect_n_stream(&mut sock_stream, 2).await;
1462
1463        //take second element, because first will be empty (signifying control frame)
1464        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); //explicitly drop the stream so we don't get warnings about not using it
1468
1469        // Ensure source actually shut down successfully.
1470        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        //1 - send READY frame (with WRONG content_type)
1623        let ready_msg = create_control_frame_with_content(
1624            ControlHeader::Ready,
1625            vec![Bytes::from(&b"test_content2"[..])],
1626        ); //can have multiple content types
1627        send_control_frame(&mut sock_sink, ready_msg).await;
1628
1629        //2 - send READY frame (with RIGHT content_type)
1630        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        //3 - wait for ACCEPT frame
1636        let mut frame_vec = collect_n_stream(&mut sock_stream, 2).await;
1637
1638        //take second element, because first will be empty (signifying control frame)
1639        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); //explicitly drop the stream so we don't get warnings about not using it
1643
1644        // Ensure source actually shut down successfully.
1645        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        //1 - send data frame (too soon!)
1658        send_data_frames(
1659            &mut sock_sink,
1660            vec![Ok(Bytes::from("bad")), Ok(Bytes::from("data"))],
1661        )
1662        .await;
1663
1664        //2 - send READY frame (with content_type)
1665        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        //3 - wait for ACCEPT frame
1671        let mut frame_vec = collect_n_stream(&mut sock_stream, 2).await;
1672
1673        //take second element, because first will be empty (signifying control frame)
1674        assert_eq!(frame_vec[0].as_ref().unwrap().len(), 0);
1675        assert_accept_frame(frame_vec[1].as_mut().unwrap(), content_type);
1676
1677        //4 - send START frame
1678        send_control_frame(&mut sock_sink, create_control_frame(ControlHeader::Start)).await;
1679
1680        //5 - send data (will go through)
1681        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        //6 - send STOP frame
1689        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); //explicitly drop the stream so we don't get warnings about not using it
1701
1702        // Ensure source actually shut down successfully.
1703        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        //1 - send START frame (with content_type)
1716        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        //4 - send data
1721        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        //5 - send STOP frame
1729        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        // Ensure source actually shut down successfully.
1741        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}