mod counters;
use counters::SinkCounters;
mod file;
mod siem;
use std::net::IpAddr;
use std::os::unix::fs::OpenOptionsExt;
use std::path::Path;
use std::sync::Arc;
use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
use std::time::Duration;
use reqwest::header::{HeaderName, HeaderValue};
use reqwest::redirect;
use serde::Serialize;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use url::Url;
use crate::StartupError;
use crate::config::Config;
pub(crate) const FIELD_CEILING_BYTES: u64 = 256 * 4 * 10 + 2;
pub(crate) const BYTES_PER_RECORD: u64 = 32 * 1024;
const _: () = assert!(
BYTES_PER_RECORD >= 3 * FIELD_CEILING_BYTES + 64 + 64 + 512 + size_of::<Decision>() as u64
);
const _: () = assert!(size_of::<Decision>() <= 256);
pub(crate) const DEFAULT_QUEUE_MAX_BYTES: u64 = 1875 * 1024 * 1024;
pub(crate) const MAX_QUEUE_MAX_BYTES: u64 = 4 * 1024 * 1024 * 1024;
pub(crate) const MAX_SAFE_CAPACITY: usize = tokio::sync::Semaphore::MAX_PERMITS;
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum CapacityError {
TooSmall,
TooLarge,
}
pub(crate) fn capacity_for(budget: u64) -> Result<usize, CapacityError> {
let records = budget / BYTES_PER_RECORD;
if records == 0 {
return Err(CapacityError::TooSmall);
}
let capacity = usize::try_from(records).map_err(|_| CapacityError::TooLarge)?;
if capacity > MAX_SAFE_CAPACITY {
return Err(CapacityError::TooLarge);
}
Ok(capacity)
}
#[derive(Clone, Debug, Serialize)]
#[serde(tag = "event", rename_all = "snake_case")]
pub(crate) enum Record {
RequestDecided(Decision),
RequestSummary(Summary),
}
#[derive(Clone, Debug, Serialize)]
pub(crate) struct Decision {
pub timestamp: String,
pub request_id: String,
pub method: String,
pub ecosystem: &'static str,
pub package: String,
pub version: String,
pub status: u16,
pub result: &'static str,
pub reason: String,
pub blocklist_revision: u64,
pub cache: &'static str,
pub duration_micros: u64,
pub bytes: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub consumer: Option<IpAddr>,
}
#[derive(Clone, Debug, Serialize)]
pub(crate) struct Summary {
pub timestamp: String,
pub requests: u64,
pub errors: u64,
pub bytes: u64,
pub mean_duration_micros: u64,
pub window_micros: u64,
pub dropped_file: u64,
pub dropped_siem: u64,
}
pub(crate) fn rfc3339(utc_micros: i64) -> String {
jiff::Timestamp::from_microsecond(utc_micros)
.map(|t| format!("{t:.6}"))
.unwrap_or_else(|_| utc_micros.to_string())
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct Drops {
pub file: u64,
pub siem: u64,
}
struct Sink {
tx: mpsc::Sender<Record>,
counters: Arc<SinkCounters>,
}
impl Sink {
fn push(&self, record: Record) {
if self.tx.try_send(record).is_err() {
self.counters.lose(1);
}
}
}
pub(crate) struct Sinks {
file: Option<Sink>,
siem: Option<Sink>,
last_change_micros: AtomicI64,
last_total: AtomicU64,
}
const SHED_QUIET: Duration = Duration::from_secs(60);
impl Sinks {
fn new(file: Option<Sink>, siem: Option<Sink>) -> Sinks {
Sinks {
file,
siem,
last_change_micros: AtomicI64::new(i64::MIN),
last_total: AtomicU64::new(0),
}
}
pub(crate) fn offer(&self, record: Record) {
match (&self.file, &self.siem) {
(Some(file), Some(siem)) => {
file.push(record.clone());
siem.push(record);
}
(Some(file), None) => file.push(record),
(None, Some(siem)) => siem.push(record),
(None, None) => {}
}
}
pub(crate) fn drops(&self) -> Drops {
Drops {
file: self
.file
.as_ref()
.map_or(0, |sink| sink.counters.take_window()),
siem: self
.siem
.as_ref()
.map_or(0, |sink| sink.counters.take_window()),
}
}
pub(crate) fn lost_total(&self) -> u64 {
[&self.file, &self.siem]
.into_iter()
.flatten()
.map(|sink| sink.counters.total())
.sum()
}
pub(crate) fn is_shedding(&self, now_utc_micros: i64) -> bool {
let total = self.lost_total();
if total != self.last_total.load(Ordering::Acquire) {
self.last_change_micros
.store(now_utc_micros, Ordering::Relaxed);
self.last_total.store(total, Ordering::Release);
}
let since = now_utc_micros.saturating_sub(self.last_change_micros.load(Ordering::Relaxed));
(0..SHED_QUIET.as_micros() as i64).contains(&since)
}
#[cfg_attr(not(feature = "test-support"), allow(dead_code))]
pub(crate) fn is_empty(&self) -> bool {
self.file.is_none() && self.siem.is_none()
}
}
const SIEM_AUTH_ENV: &str = "PROBATION_SIEM_AUTH";
const SIEM_AUTH_REJECTED: &str = "PROBATION_SIEM_AUTH is not a valid HTTP header value; \
it must be printable ASCII with no line break";
const SIEM_REQUEST_TIMEOUT: Duration = Duration::from_secs(3);
const SIEM_CONNECT_TIMEOUT: Duration = Duration::from_secs(1);
pub(crate) fn build(
config: &Config,
drain: CancellationToken,
) -> Result<(Sinks, Vec<JoinHandle<()>>), StartupError> {
let mut tasks = Vec::new();
if let Some(path) = &config.log_file_path {
probe(path, config.log_consumer_identification)?;
}
let file = config.log_file_path.as_ref().map(|path| {
let (tx, rx) = mpsc::channel(
capacity_for(config.log_queue_max_bytes.get()).expect("config validated"),
);
let counters = Arc::new(SinkCounters::new());
tasks.push(tokio::spawn(file::run(
path.clone(),
config.log_file_max_bytes,
rx,
drain.clone(),
Arc::clone(&counters),
)));
Sink { tx, counters }
});
let mut siem_auth_header = None;
let siem = match &config.siem_url {
Some(url) => {
let auth = siem_auth(&config.siem_auth_header)?;
siem_auth_header = auth.as_ref().map(|(name, _)| name.as_str().to_owned());
let client = reqwest::Client::builder()
.redirect(redirect::Policy::none())
.timeout(SIEM_REQUEST_TIMEOUT)
.connect_timeout(SIEM_CONNECT_TIMEOUT)
.no_proxy()
.build()
.map_err(|err| {
tracing::error!(error = %err, "the SIEM delivery client could not be built");
StartupError::Delivery("the SIEM delivery client could not be built")
})?;
let (tx, rx) = mpsc::channel(
capacity_for(config.siem_queue_max_bytes.get()).expect("config validated"),
);
let counters = Arc::new(SinkCounters::new());
tasks.push(tokio::spawn(siem::run(
client,
url.clone(),
auth,
rx,
drain.clone(),
Arc::clone(&counters),
)));
Some(Sink { tx, counters })
}
None => None,
};
if config.log_file_path.is_some() || config.siem_url.is_some() {
tracing::info!(
log_file = ?config.log_file_path,
log_file_max_bytes = config.log_file_max_bytes.get(),
siem_url = config.siem_url.as_ref().map(Url::as_str),
siem_auth_header = siem_auth_header.as_deref(),
"decision records are delivered off this process"
);
}
Ok((Sinks::new(file, siem), tasks))
}
const LOG_FILE_REJECTED: &str = "log_file_path cannot be opened for append, and \
log_consumer_identification is on: peer addresses would be collected with nowhere \
durable to go";
fn probe(path: &Path, consumer_identification: bool) -> Result<(), StartupError> {
let Err(err) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.mode(0o600)
.open(path)
else {
return Ok(());
};
if consumer_identification {
return Err(StartupError::Delivery(LOG_FILE_REJECTED));
}
tracing::error!(
key = "log_file_path",
path = %path.display(),
error = %err,
"log_file_path cannot be opened for append; decision records will be dropped and \
counted until it can"
);
Ok(())
}
fn siem_auth(name: &HeaderName) -> Result<Option<(HeaderName, HeaderValue)>, StartupError> {
let value = match std::env::var(SIEM_AUTH_ENV) {
Ok(value) => value,
Err(std::env::VarError::NotPresent) => return Ok(None),
Err(std::env::VarError::NotUnicode(_)) => {
return Err(StartupError::Delivery(SIEM_AUTH_REJECTED));
}
};
let mut value =
HeaderValue::from_str(&value).map_err(|_| StartupError::Delivery(SIEM_AUTH_REJECTED))?;
value.set_sensitive(true);
Ok(Some((name.clone(), value)))
}
#[cfg(test)]
mod tests {
use super::*;
fn max_safe_capacity_u64() -> u64 {
u64::try_from(MAX_SAFE_CAPACITY).unwrap_or(u64::MAX)
}
proptest::proptest! {
#[test]
fn rl6_no_budget_panics_or_yields_a_bad_capacity(budget: u64) {
let records = budget / BYTES_PER_RECORD;
match capacity_for(budget) {
Ok(capacity) => {
proptest::prop_assert!(
(1..=MAX_SAFE_CAPACITY).contains(&capacity),
"budget {} yielded capacity {}, which mpsc::channel would refuse",
budget,
capacity
);
proptest::prop_assert_eq!(u64::try_from(capacity).unwrap(), records);
}
Err(CapacityError::TooSmall) => {
proptest::prop_assert_eq!(records, 0, "budget {} holds a whole record", budget);
}
Err(CapacityError::TooLarge) => {
proptest::prop_assert!(
records > max_safe_capacity_u64(),
"budget {} is within the channel's limit and was still refused",
budget
);
}
}
}
}
fn summary() -> Record {
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,
})
}
fn full_file_sink() -> (Sinks, mpsc::Receiver<Record>) {
let (tx, rx) = mpsc::channel(1);
let sink = Sink {
tx,
counters: Arc::new(SinkCounters::new()),
};
(Sinks::new(Some(sink), None), rx)
}
#[test]
fn rl7_queue_full_counts_in_the_window() {
let (sinks, _rx) = full_file_sink();
sinks.offer(summary());
assert_eq!(sinks.drops().file, 0, "the first record fits");
sinks.offer(summary());
assert_eq!(sinks.drops().file, 1);
}
#[test]
fn rl8_the_second_window_reports_only_its_own_losses() {
let (sinks, _rx) = full_file_sink();
for _ in 0..4 {
sinks.offer(summary()); }
assert_eq!(sinks.drops().file, 3);
sinks.offer(summary());
assert_eq!(sinks.drops().file, 1, "not 4: the first window was reset");
assert_eq!(sinks.drops().file, 0);
assert_eq!(sinks.lost_total(), 4);
}
}