use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use reqwest::header::{CONTENT_TYPE, HeaderName, HeaderValue};
use reqwest::{Client, StatusCode};
use tokio::sync::mpsc;
use tokio::time::Instant;
use tokio_util::sync::CancellationToken;
use url::Url;
use super::Record;
use super::counters::SinkCounters;
const BATCH_RECORDS: usize = 256;
const BATCH_INTERVAL: Duration = Duration::from_secs(2);
const BACKOFF: [Duration; 3] = [
Duration::from_millis(100),
Duration::from_millis(500),
Duration::from_secs(2),
];
const DRAIN_DEADLINE: Duration = Duration::from_secs(5);
pub(super) async fn run(
client: Client,
url: Url,
auth: Option<(HeaderName, HeaderValue)>,
mut rx: mpsc::Receiver<Record>,
drain: CancellationToken,
counters: Arc<SinkCounters>,
) {
let mut batch: Vec<Record> = Vec::new();
let mut deadline = Instant::now();
loop {
tokio::select! {
received = rx.recv() => match received {
Some(record) => {
if batch.is_empty() {
deadline = Instant::now() + BATCH_INTERVAL;
}
batch.push(record);
if batch.len() >= BATCH_RECORDS {
send(&client, &url, auth.as_ref(), &mut batch, &counters, &drain).await;
}
}
None => break,
},
() = tokio::time::sleep_until(deadline), if !batch.is_empty() => {
send(&client, &url, auth.as_ref(), &mut batch, &counters, &drain).await;
}
() = drain.cancelled() => {
let _ = tokio::time::timeout(DRAIN_DEADLINE, async {
while let Ok(record) = rx.try_recv() {
batch.push(record);
if batch.len() >= BATCH_RECORDS {
send(&client, &url, auth.as_ref(), &mut batch, &counters, &drain).await;
}
}
if !batch.is_empty() {
send(&client, &url, auth.as_ref(), &mut batch, &counters, &drain).await;
}
})
.await;
let lost = (batch.len() + rx.len()) as u64;
if lost != 0 {
counters.lose(lost);
}
break;
}
}
}
}
async fn send(
client: &Client,
url: &Url,
auth: Option<&(HeaderName, HeaderValue)>,
batch: &mut Vec<Record>,
counters: &SinkCounters,
drain: &CancellationToken,
) {
let count = batch.len() as u64;
let mut body = String::new();
for record in batch.drain(..) {
if let Ok(line) = serde_json::to_string(&record) {
body.push_str(&line);
body.push('\n');
} else {
counters.lose(1);
}
}
let mut unsent = Unsent { count, counters };
let body = Bytes::from(body);
let mut attempt = 0;
loop {
let mut request = client
.post(url.clone())
.header(CONTENT_TYPE, "application/x-ndjson")
.body(body.clone());
if let Some((name, value)) = auth {
request = request.header(name.clone(), value.clone());
}
let retryable = match request.send().await {
Ok(response) if response.status().is_success() => {
unsent.delivered();
return;
}
Ok(response) => {
let status = response.status();
tracing::warn!(
status = status.as_u16(),
"the SIEM collector refused a batch of decision records"
);
status.is_server_error() || status == StatusCode::TOO_MANY_REQUESTS
}
Err(err) => {
tracing::warn!(
error = %err,
"a batch of decision records could not be sent to the SIEM collector"
);
true
}
};
if !retryable {
break;
}
let backoff = BACKOFF[attempt.min(BACKOFF.len() - 1)];
tokio::select! {
() = tokio::time::sleep(backoff) => {}
() = drain.cancelled() => return,
}
attempt += 1;
}
tracing::warn!(
records = count,
"a batch of decision records was dropped after the SIEM collector refused it"
);
}
struct Unsent<'a> {
count: u64,
counters: &'a SinkCounters,
}
impl Unsent<'_> {
fn delivered(&mut self) {
self.count = 0;
}
}
impl Drop for Unsent<'_> {
fn drop(&mut self) {
if self.count != 0 {
self.counters.lose(self.count);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn rl7_siem_drain_deadline_and_unsent_count_in_the_window() {
use crate::delivery::Summary;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = Url::parse(&format!("http://{}/", listener.local_addr().unwrap())).unwrap();
let hold = tokio::spawn(async move {
let mut held = Vec::new();
while let Ok((stream, _)) = listener.accept().await {
held.push(stream);
}
});
let total = BATCH_RECORDS + 44;
let (tx, rx) = mpsc::channel(total);
for _ in 0..total {
tx.try_send(Record::RequestSummary(Summary {
timestamp: String::new(),
requests: 0,
errors: 0,
bytes: 0,
mean_duration_micros: 0,
window_micros: 0,
dropped_file: 0,
dropped_siem: 0,
}))
.unwrap();
}
let drain = CancellationToken::new();
drain.cancel();
let counters = Arc::new(SinkCounters::new());
run(Client::new(), url, None, rx, drain, Arc::clone(&counters)).await;
hold.abort();
assert_eq!(counters.take_window(), total as u64);
}
#[test]
fn rl7_unsent_counts_an_undelivered_batch() {
let counters = SinkCounters::new();
drop(Unsent {
count: 5,
counters: &counters,
});
let mut delivered = Unsent {
count: 7,
counters: &counters,
};
delivered.delivered();
drop(delivered);
assert_eq!(counters.take_window(), 5);
}
}