#![allow(dead_code)]
use std::collections::HashMap;
use std::future::Future;
use std::io::Write;
use std::net::SocketAddr;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicI64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant, SystemTime};
use async_trait::async_trait;
use probation::clock::{Clock, SystemClock};
use probation::config::Config;
use probation::upstream::{
ArtifactBody, ArtifactRequest, MetadataRequest, MetadataResponse, OriginSet, ReqwestTransport,
Transport, UpstreamError, UpstreamValidators,
};
use probation::{App, AppDeps, Running};
use tempfile::TempDir;
use url::Url;
pub fn sample_config() -> Config {
let mut config = Config::load(Path::new("config.sample.toml"))
.expect("config.sample.toml is valid configuration");
config.listen = SocketAddr::from(([127, 0, 0, 1], 0));
config
}
pub struct TestClock {
start_utc_micros: i64,
utc_micros: AtomicI64,
monotonic_base: Instant,
monotonic_carry_micros: AtomicI64,
}
impl TestClock {
pub fn at(utc_micros: i64) -> Arc<TestClock> {
Arc::new(TestClock {
start_utc_micros: utc_micros,
utc_micros: AtomicI64::new(utc_micros),
monotonic_base: Instant::now(),
monotonic_carry_micros: AtomicI64::new(0),
})
}
pub fn at_rfc3339(text: &str) -> Arc<TestClock> {
TestClock::at(parse_rfc3339(text))
}
pub fn set_rfc3339(&self, text: &str) {
self.utc_micros.store(parse_rfc3339(text), Ordering::SeqCst);
}
pub fn advance_seconds(&self, seconds: i64) {
self.utc_micros
.fetch_add(seconds * 1_000_000, Ordering::SeqCst);
}
pub fn rewind_wall_clock_seconds(&self, seconds: i64) {
self.monotonic_carry_micros
.fetch_add(seconds * 1_000_000, Ordering::SeqCst);
self.utc_micros
.fetch_sub(seconds * 1_000_000, Ordering::SeqCst);
}
pub fn shared(self: &Arc<Self>) -> Arc<dyn Clock> {
let clock: Arc<TestClock> = Arc::clone(self);
clock
}
}
impl Clock for TestClock {
fn now_utc_micros(&self) -> i64 {
self.utc_micros.load(Ordering::SeqCst)
}
fn now_monotonic(&self) -> Instant {
let elapsed = self
.now_utc_micros()
.saturating_sub(self.start_utc_micros)
.max(0)
.saturating_add(self.monotonic_carry_micros.load(Ordering::SeqCst));
self.monotonic_base + Duration::from_micros(elapsed as u64)
}
}
pub fn parse_rfc3339(text: &str) -> i64 {
text.parse::<jiff::Timestamp>()
.expect("a test timestamp")
.as_microsecond()
}
pub struct TestServer {
running: Running,
client: reqwest::Client,
_data_dir: Option<TempDir>,
}
impl TestServer {
pub async fn start() -> TestServer {
TestServer::start_with(sample_config(), Arc::new(SystemClock)).await
}
pub async fn start_with(config: Config, clock: Arc<dyn Clock>) -> TestServer {
let (transport, origins) = fake_upstream();
TestServer::start_with_upstream(config, clock, transport, origins).await
}
pub async fn start_with_upstream(
mut config: Config,
clock: Arc<dyn Clock>,
transport: Arc<dyn Transport>,
origins: OriginSet,
) -> TestServer {
let data_dir = tempfile::tempdir().expect("a temporary data directory");
config.data_dir = data_dir.path().to_path_buf();
let running = start(
config,
clock,
transport,
origins,
probation::osv::unreachable_client(),
None,
)
.await;
TestServer {
running,
client: downstream_client(),
_data_dir: Some(data_dir),
}
}
pub async fn start_with_upstream_and_osv(
mut config: Config,
clock: Arc<dyn Clock>,
transport: Arc<dyn Transport>,
origins: OriginSet,
osv_client: reqwest::Client,
osv_base_url: Url,
) -> TestServer {
let data_dir = tempfile::tempdir().expect("a temporary data directory");
config.data_dir = data_dir.path().to_path_buf();
let running = start(
config,
clock,
transport,
origins,
osv_client,
Some(osv_base_url),
)
.await;
TestServer {
running,
client: downstream_client(),
_data_dir: Some(data_dir),
}
}
pub async fn start_in(
data_dir: &Path,
mut config: Config,
clock: Arc<dyn Clock>,
) -> TestServer {
config.data_dir = data_dir.to_path_buf();
let (transport, origins) = fake_upstream();
let running = start(
config,
clock,
transport,
origins,
probation::osv::unreachable_client(),
None,
)
.await;
TestServer {
running,
client: downstream_client(),
_data_dir: None,
}
}
pub async fn start_in_with_registry(
data_dir: &Path,
mut config: Config,
clock: Arc<dyn Clock>,
registry: Arc<FakeRegistry>,
) -> TestServer {
config.data_dir = data_dir.to_path_buf();
let running = start(
config,
clock,
registry,
fake_origins(),
probation::osv::unreachable_client(),
None,
)
.await;
TestServer {
running,
client: downstream_client(),
_data_dir: None,
}
}
pub fn running(&self) -> &Running {
&self.running
}
pub fn store_commands(&self) -> u64 {
self.running.app().store().commands_issued()
}
pub async fn json(&self, path: &str) -> serde_json::Value {
let response = self.get(path).await;
let status = response.status();
let body = response.text().await.expect("a body");
assert!(
status.is_success(),
"{path} answered {status}, not a document: {body}"
);
serde_json::from_str(&body)
.unwrap_or_else(|err| panic!("{path} is not JSON: {err}: {body}"))
}
pub fn url(&self, path: &str) -> String {
format!("http://{}{path}", self.running.local_addr)
}
pub fn local_addr(&self) -> SocketAddr {
self.running.local_addr
}
pub fn app(&self) -> Arc<probation::App> {
Arc::clone(self.running.app())
}
pub async fn get(&self, path: &str) -> reqwest::Response {
self.client
.get(self.url(path))
.send()
.await
.expect("the request completes")
}
pub async fn get_with_headers(
&self,
path: &str,
headers: &[(&str, &str)],
) -> reqwest::Response {
let mut request = self.client.get(self.url(path));
for (name, value) in headers {
request = request.header(*name, *value);
}
request.send().await.expect("the request completes")
}
pub async fn post(&self, path: &str) -> reqwest::Response {
self.client
.post(self.url(path))
.send()
.await
.expect("the request completes")
}
pub async fn head(&self, path: &str) -> reqwest::Response {
self.client
.head(self.url(path))
.send()
.await
.expect("the request completes")
}
pub async fn status(&self, path: &str) -> u16 {
self.get(path).await.status().as_u16()
}
pub async fn raw_get(&self, raw_path: &str, headers: &[(&str, &str)]) -> String {
let addr = self.running.local_addr;
let mut request =
format!("GET {raw_path} HTTP/1.1\r\nHost: {addr}\r\nConnection: close\r\n");
for (name, value) in headers {
request.push_str(&format!("{name}: {value}\r\n"));
}
request.push_str("\r\n");
self.raw_send(&request).await
}
pub async fn raw_send(&self, raw_request: &str) -> String {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let addr = self.running.local_addr;
let mut stream = tokio::net::TcpStream::connect(addr)
.await
.expect("the server accepts a connection");
stream
.write_all(raw_request.as_bytes())
.await
.expect("the request is written");
let mut response = Vec::new();
stream
.read_to_end(&mut response)
.await
.expect("the response is read");
String::from_utf8_lossy(&response).into_owned()
}
pub async fn raw_get_status(&self, raw_path: &str) -> u16 {
let response = self.raw_get(raw_path, &[]).await;
let status = response
.split_whitespace()
.nth(1)
.unwrap_or_else(|| panic!("a status line, got {response:?}"));
status.parse().expect("a numeric status")
}
pub async fn wait_for_status(&self, path: &str, expected: u16, within: Duration) {
let deadline = Instant::now() + within;
let mut last = self.status(path).await;
while last != expected {
if Instant::now() >= deadline {
panic!("{path} answered {last}, not {expected}, within {within:?}");
}
tokio::time::sleep(Duration::from_millis(25)).await;
last = self.status(path).await;
}
}
pub async fn shutdown(self) {
self.running.shutdown().await.expect("a clean shutdown");
}
}
pub fn downstream_client() -> reqwest::Client {
reqwest::Client::builder()
.no_proxy()
.build()
.expect("a test HTTP client")
}
pub async fn wait_until(what: &str, within: Duration, mut condition: impl FnMut() -> bool) {
let deadline = Instant::now() + within;
while !condition() {
assert!(
Instant::now() < deadline,
"{what} did not happen within {within:?}"
);
tokio::task::yield_now().await;
tokio::time::sleep(Duration::from_millis(2)).await;
}
}
pub struct SlowReader {
stream: tokio::net::TcpStream,
}
impl SlowReader {
pub async fn get(server: &TestServer, path: &str) -> SlowReader {
use tokio::io::AsyncWriteExt;
let addr = server.local_addr();
let mut stream = tokio::net::TcpStream::connect(addr)
.await
.expect("the server accepts a connection");
let request = format!("GET {path} HTTP/1.1\r\nHost: {addr}\r\nConnection: close\r\n\r\n");
stream
.write_all(request.as_bytes())
.await
.expect("the request is written");
SlowReader { stream }
}
pub async fn read_what_arrived(mut self) -> Vec<u8> {
use tokio::io::AsyncReadExt;
let mut response = Vec::new();
let _ = self.stream.read_to_end(&mut response).await;
response
}
}
pub fn raw_status(response: &str) -> u16 {
response
.split_whitespace()
.nth(1)
.unwrap_or_else(|| panic!("a status line, got {response:?}"))
.parse()
.expect("a numeric status")
}
pub fn raw_header(response: &str, name: &str) -> Option<String> {
response
.split("\r\n\r\n")
.next()?
.lines()
.skip(1)
.find_map(|line| {
let (header, value) = line.split_once(':')?;
header
.trim()
.eq_ignore_ascii_case(name)
.then(|| value.trim().to_owned())
})
}
pub fn body_bytes(response: &[u8]) -> &[u8] {
match response.windows(4).position(|window| window == b"\r\n\r\n") {
Some(at) => &response[at + 4..],
None => &[],
}
}
async fn start(
config: Config,
clock: Arc<dyn Clock>,
transport: Arc<dyn Transport>,
origins: OriginSet,
osv_client: reqwest::Client,
osv_base_url: Option<Url>,
) -> Running {
App::start(AppDeps {
config,
clock,
transport,
origins,
osv_client,
osv_base_url,
})
.await
.expect("the server binds and starts")
}
#[derive(Clone, Debug)]
pub enum FakeAnswer {
Body(String),
Missing,
Fail(UpstreamError),
Validated {
etag: String,
body: String,
},
AlwaysNotModified {
etag: Option<String>,
last_modified: Option<String>,
},
Truncated {
declared_length: u64,
body: String,
},
Gated {
head: String,
tail: String,
gate: Arc<Gate>,
},
PanicsMidStream {
head: String,
},
GatedMetadata {
body: String,
gate: Arc<Gate>,
},
GatedNotModified {
etag: String,
gate: Arc<Gate>,
},
PanicsMidRefresh {
gate: Arc<Gate>,
},
}
#[derive(Debug)]
pub struct Gate {
reached: tokio::sync::Semaphore,
released: tokio::sync::Semaphore,
}
impl Gate {
pub fn new() -> Arc<Gate> {
Arc::new(Gate {
reached: tokio::sync::Semaphore::new(0),
released: tokio::sync::Semaphore::new(0),
})
}
pub async fn wait_until_reached(&self) {
let _ = self.reached.acquire().await.expect("the gate is open");
}
pub fn release(&self) {
self.released.add_permits(1);
}
async fn arrive(&self) {
self.reached.add_permits(1);
let _ = self.released.acquire().await.expect("the gate is open");
}
}
#[derive(Default)]
pub struct FakeRegistry {
answers: Mutex<HashMap<String, FakeAnswer>>,
calls: Mutex<Vec<Url>>,
conditional: Mutex<Vec<(String, UpstreamValidators)>>,
not_modified: std::sync::atomic::AtomicUsize,
offline: std::sync::atomic::AtomicBool,
}
impl FakeRegistry {
pub fn new() -> Arc<FakeRegistry> {
Arc::new(FakeRegistry::default())
}
pub fn answer(self: &Arc<Self>, path: &str, answer: FakeAnswer) -> Arc<FakeRegistry> {
self.answers
.lock()
.expect("the fake registry's answers")
.insert(path.to_owned(), answer);
Arc::clone(self)
}
pub fn calls(&self) -> Vec<Url> {
self.calls
.lock()
.expect("the fake registry's calls")
.clone()
}
pub fn conditional_calls(&self, path: &str) -> Vec<UpstreamValidators> {
self.conditional
.lock()
.expect("the fake registry's conditional requests")
.iter()
.filter(|(asked, _)| asked == path)
.map(|(_, validators)| validators.clone())
.collect()
}
pub fn not_modified_answers(&self) -> usize {
self.not_modified.load(std::sync::atomic::Ordering::SeqCst)
}
pub fn go_offline(&self) {
self.offline
.store(true, std::sync::atomic::Ordering::SeqCst);
}
fn record(&self, url: &Url) -> FakeAnswer {
self.calls
.lock()
.expect("the fake registry's calls")
.push(url.clone());
if self.offline.load(std::sync::atomic::Ordering::SeqCst) {
return FakeAnswer::Fail(UpstreamError::Transport(
"the upstream registry is unreachable".to_owned(),
));
}
self.answers
.lock()
.expect("the fake registry's answers")
.get(url.path())
.cloned()
.unwrap_or(FakeAnswer::Missing)
}
}
#[async_trait]
impl Transport for FakeRegistry {
async fn fetch_metadata(
&self,
req: MetadataRequest,
) -> Result<MetadataResponse, UpstreamError> {
if let Some(validators) = req.validators.as_ref().filter(|v| !v.is_empty()) {
self.conditional
.lock()
.expect("the fake registry's conditional requests")
.push((req.url.path().to_owned(), validators.clone()));
}
match self.record(&req.url) {
FakeAnswer::Body(body) => {
if body.len() as u64 > req.max_bytes {
return Err(UpstreamError::TooLarge {
limit: req.max_bytes,
});
}
Ok(MetadataResponse::Fresh {
body: body.into(),
validators: UpstreamValidators::default(),
})
}
FakeAnswer::Validated { etag, body } => {
let validators = UpstreamValidators {
etag: Some(etag.clone()),
last_modified: None,
};
if req.validators.and_then(|sent| sent.etag).as_deref() == Some(etag.as_str()) {
self.not_modified
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
return Ok(MetadataResponse::NotModified { validators });
}
if body.len() as u64 > req.max_bytes {
return Err(UpstreamError::TooLarge {
limit: req.max_bytes,
});
}
Ok(MetadataResponse::Fresh {
body: body.into(),
validators,
})
}
FakeAnswer::AlwaysNotModified {
etag,
last_modified,
} => {
self.not_modified
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(MetadataResponse::NotModified {
validators: UpstreamValidators {
etag,
last_modified,
},
})
}
FakeAnswer::Missing => Ok(MetadataResponse::Missing),
FakeAnswer::Fail(err) => Err(err),
FakeAnswer::Truncated { body, .. } | FakeAnswer::Gated { head: body, .. } => {
Ok(MetadataResponse::Fresh {
body: body.into(),
validators: UpstreamValidators::default(),
})
}
FakeAnswer::PanicsMidStream { .. } => Ok(MetadataResponse::Missing),
FakeAnswer::GatedMetadata { body, gate } => {
gate.arrive().await;
Ok(MetadataResponse::Fresh {
body: body.into(),
validators: UpstreamValidators::default(),
})
}
FakeAnswer::GatedNotModified { etag, gate } => {
gate.arrive().await;
self.not_modified
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(MetadataResponse::NotModified {
validators: UpstreamValidators {
etag: Some(etag),
last_modified: None,
},
})
}
FakeAnswer::PanicsMidRefresh { gate } => {
gate.arrive().await;
panic!("the upstream metadata refresh panicked mid-flight");
}
}
}
async fn open_artifact(&self, req: ArtifactRequest) -> Result<ArtifactBody, UpstreamError> {
let capped = |body: String, declared: Option<u64>| ArtifactBody {
declared_length: declared,
stream: probation::upstream::capped(
Box::pin(futures_util::stream::once(async move {
Ok(bytes::Bytes::from(body))
})),
req.max_bytes,
),
};
match self.record(&req.url) {
FakeAnswer::Body(body) => {
let declared = body.len() as u64;
Ok(capped(body, Some(declared)))
}
FakeAnswer::Truncated {
declared_length,
body,
} => Ok(capped(body, Some(declared_length))),
FakeAnswer::Gated { head, tail, gate } => {
let declared = (head.len() + tail.len()) as u64;
let stream = futures_util::stream::unfold(
(Some(head), Some(tail), gate),
|(head, tail, gate)| async move {
if let Some(head) = head {
return Some((Ok(bytes::Bytes::from(head)), (None, tail, gate)));
}
let tail = tail?;
gate.arrive().await;
Some((Ok(bytes::Bytes::from(tail)), (None, None, gate)))
},
);
Ok(ArtifactBody {
declared_length: Some(declared),
stream: probation::upstream::capped(Box::pin(stream), req.max_bytes),
})
}
FakeAnswer::PanicsMidStream { head } => {
let declared = (head.len() + 1) as u64;
let stream = futures_util::stream::unfold(Some(head), |head| async move {
match head {
Some(head) => Some((Ok(bytes::Bytes::from(head)), None)),
None => panic!("the upstream transfer panicked mid-body"),
}
});
Ok(ArtifactBody {
declared_length: Some(declared),
stream: probation::upstream::capped(Box::pin(stream), req.max_bytes),
})
}
FakeAnswer::Missing => Err(UpstreamError::Status(404)),
FakeAnswer::Fail(err) => Err(err),
FakeAnswer::Validated { .. }
| FakeAnswer::AlwaysNotModified { .. }
| FakeAnswer::GatedMetadata { .. }
| FakeAnswer::GatedNotModified { .. }
| FakeAnswer::PanicsMidRefresh { .. } => Err(UpstreamError::Status(404)),
}
}
}
pub fn fake_origins() -> OriginSet {
OriginSet::for_tests(
Url::parse("https://npm.invalid").expect("a fake npm origin"),
Url::parse("https://pypi.invalid").expect("a fake pypi origin"),
Url::parse("https://files.invalid").expect("a fake artifact origin"),
)
}
pub fn fake_upstream() -> (Arc<dyn Transport>, OriginSet) {
(FakeRegistry::new(), fake_origins())
}
pub fn npm_upstream_path(name: &str) -> String {
fake_origins()
.url_for(probation::upstream::OriginKind::NpmMetadata, &[name])
.expect("the fake origin builds a URL for this name")
.path()
.to_owned()
}
pub fn pypi_upstream_path(name: &str) -> String {
fake_origins()
.url_for(
probation::upstream::OriginKind::PypiMetadata,
&["simple", name, ""],
)
.expect("the fake origin builds a URL for this project")
.path()
.to_owned()
}
pub async fn body_error(response: reqwest::Response) -> String {
let body = response.text().await.expect("a body");
let value: serde_json::Value = serde_json::from_str(&body)
.unwrap_or_else(|err| panic!("the body is not JSON: {err}: {body}"));
value["error"]
.as_str()
.unwrap_or_else(|| panic!("the error body names an error: {body}"))
.to_owned()
}
pub fn fixture(relative: &str) -> String {
let path = Path::new("tests/fixtures").join(relative);
std::fs::read_to_string(&path)
.unwrap_or_else(|err| panic!("the fixture {} is readable: {err}", path.display()))
}
pub struct WiremockUpstream {
pub server: wiremock::MockServer,
pub origins: OriginSet,
pub transport: Arc<dyn Transport>,
}
impl WiremockUpstream {
pub async fn start() -> WiremockUpstream {
let server = wiremock::MockServer::start().await;
let origin = Url::parse(&server.uri()).expect("the wiremock origin is a URL");
let origins = OriginSet::for_tests(origin.clone(), origin.clone(), origin);
let transport = Arc::new(ReqwestTransport::for_origins(origins.clone()));
WiremockUpstream {
server,
origins,
transport,
}
}
pub fn url(&self, path: &str) -> String {
format!("{}{path}", self.server.uri())
}
}
pub fn config_with_open_blocklist(dir: &Path) -> Config {
let blocklist_file = dir.join("blocklist.json");
std::fs::write(
&blocklist_file,
snapshot(1, "2020-01-01T00:00:00Z", "2099-01-01T00:00:00Z", ""),
)
.expect("the blocklist is written");
let mut config = sample_config();
config.blocklist_file = blocklist_file;
config
}
pub fn run_and_kill<F>(body: impl FnOnce() -> F)
where
F: Future<Output = ()>,
{
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()
.expect("a test runtime");
runtime.block_on(body());
runtime.shutdown_timeout(Duration::ZERO);
}
pub fn run<T>(body: impl Future<Output = T>) -> T {
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()
.expect("a test runtime");
runtime.block_on(body)
}
pub mod logs {
use std::io;
use std::sync::{Arc, Mutex, OnceLock};
#[derive(Clone, Default)]
pub struct Captured(Arc<Mutex<Vec<u8>>>);
impl io::Write for Captured {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.0
.lock()
.expect("the capture buffer")
.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl<'a> tracing_subscriber::fmt::MakeWriter<'a> for Captured {
type Writer = Captured;
fn make_writer(&'a self) -> Captured {
self.clone()
}
}
pub fn capture() -> Captured {
install(tracing::Level::ERROR)
}
pub fn capture_warn() -> Captured {
install(tracing::Level::WARN)
}
pub fn capture_info() -> Captured {
install(tracing::Level::INFO)
}
fn install(level: tracing::Level) -> Captured {
static CAPTURED: OnceLock<Captured> = OnceLock::new();
CAPTURED
.get_or_init(|| {
let captured = Captured::default();
let subscriber = tracing_subscriber::fmt()
.with_writer(captured.clone())
.with_ansi(false)
.with_max_level(level)
.finish();
let _ = tracing::subscriber::set_global_default(subscriber);
captured
})
.clone()
}
impl Captured {
pub fn lines_mentioning(&self, needle: &str) -> usize {
self.lines_containing_all(&[needle])
}
pub fn lines_containing_all(&self, needles: &[&str]) -> usize {
String::from_utf8_lossy(&self.0.lock().expect("the capture buffer"))
.lines()
.filter(|line| needles.iter().all(|needle| line.contains(needle)))
.count()
}
pub fn text(&self) -> String {
String::from_utf8_lossy(&self.0.lock().expect("the capture buffer")).into_owned()
}
}
}
pub fn snapshot(revision: u64, generated_at: &str, expires_at: &str, blocked: &str) -> String {
format!(
r#"{{"schema_version":1,"revision":{revision},"generated_at":"{generated_at}","expires_at":"{expires_at}","blocked_packages":[{blocked}],"blocked_hashes":[]}}"#
)
}
pub fn snapshot_with(
revision: u64,
generated_at: &str,
expires_at: &str,
blocked: &str,
hashes: &str,
) -> String {
format!(
r#"{{"schema_version":1,"revision":{revision},"generated_at":"{generated_at}","expires_at":"{expires_at}","blocked_packages":[{blocked}],"blocked_hashes":[{hashes}]}}"#
)
}
pub fn publish_blocklist(server: &TestServer, document: &str, now_micros: i64) {
let snapshot = probation::policy::BlocklistSnapshot::parse_and_validate(
document.as_bytes(),
now_micros,
)
.expect("the test blocklist is valid");
server
.running()
.app()
.publish_blocklist(std::sync::Arc::new(snapshot));
}
pub struct TriggerClock {
inner: Arc<TestClock>,
armed: std::sync::atomic::AtomicBool,
app: Mutex<Option<std::sync::Weak<probation::App>>>,
document: Mutex<Option<String>>,
}
impl TriggerClock {
pub fn at_rfc3339(text: &str) -> Arc<TriggerClock> {
Arc::new(TriggerClock {
inner: TestClock::at_rfc3339(text),
armed: std::sync::atomic::AtomicBool::new(false),
app: Mutex::new(None),
document: Mutex::new(None),
})
}
pub fn shared(self: &Arc<Self>) -> Arc<dyn Clock> {
let clock: Arc<TriggerClock> = Arc::clone(self);
clock
}
pub fn arm(&self, server: &TestServer, document: &str) {
*self.app.lock().expect("the trigger's app") = Some(Arc::downgrade(server.running().app()));
*self.document.lock().expect("the trigger's document") = Some(document.to_owned());
self.armed.store(true, std::sync::atomic::Ordering::SeqCst);
}
pub fn fired(&self) -> bool {
!self.armed.load(std::sync::atomic::Ordering::SeqCst)
&& self
.document
.lock()
.expect("the trigger's document")
.is_some()
}
}
impl Clock for TriggerClock {
fn now_utc_micros(&self) -> i64 {
let now = self.inner.now_utc_micros();
if self.armed.swap(false, std::sync::atomic::Ordering::SeqCst)
&& let (Some(app), Some(document)) = (
self.app
.lock()
.expect("the trigger's app")
.as_ref()
.and_then(std::sync::Weak::upgrade),
self.document
.lock()
.expect("the trigger's document")
.clone(),
)
{
let snapshot = probation::policy::BlocklistSnapshot::parse_and_validate(
document.as_bytes(),
now,
)
.expect("the test blocklist is valid");
app.publish_blocklist(std::sync::Arc::new(snapshot));
}
now
}
fn now_monotonic(&self) -> Instant {
self.inner.now_monotonic()
}
}
pub fn npm_artifact_upstream_path(name: &str, filename: &str) -> String {
format!("/{name}/-/{filename}")
}
pub fn npm_tarball_url(name: &str, filename: &str) -> String {
format!("https://npm.invalid/{name}/-/{filename}")
}
pub fn artifact_path(document: &serde_json::Value, version: &str) -> String {
let tarball = document["versions"][version]["dist"]["tarball"]
.as_str()
.unwrap_or_else(|| panic!("version {version} has a rewritten tarball URL: {document}"));
let url = Url::parse(tarball).expect("the rewritten tarball URL parses");
url.path().to_owned()
}
pub fn replace_atomically(path: &Path, contents: &str) {
let temp = path.with_extension("next");
std::fs::write(&temp, contents).expect("the replacement is written");
std::fs::rename(&temp, path).expect("the replacement is renamed into place");
}
pub fn replace_atomically_preserving_mtime(path: &Path, contents: &str) {
let modified = modified_of(path);
let temp = path.with_extension("next");
std::fs::write(&temp, contents).expect("the replacement is written");
set_modified(&temp, modified);
std::fs::rename(&temp, path).expect("the replacement is renamed into place");
}
pub fn rewrite_in_place(path: &Path, contents: &str) {
let mut file = std::fs::OpenOptions::new()
.write(true)
.truncate(true)
.open(path)
.expect("the existing file is opened");
file.write_all(contents.as_bytes())
.expect("the rewrite is written");
file.sync_all().expect("the rewrite reaches the filesystem");
}
pub fn set_modified(path: &Path, modified: SystemTime) {
let file = std::fs::OpenOptions::new()
.write(true)
.open(path)
.expect("the file is opened to set its times");
file.set_modified(modified)
.expect("the modification time is set");
}
pub fn modified_of(path: &Path) -> SystemTime {
std::fs::metadata(path)
.expect("the file exists")
.modified()
.expect("the filesystem records modification times")
}
pub fn database_path(data_dir: &Path) -> PathBuf {
data_dir.join("state").join("firewall.db")
}
pub fn wal_path(data_dir: &Path) -> PathBuf {
data_dir.join("state").join("firewall.db-wal")
}
#[derive(Debug)]
pub struct RatedLoadResult {
pub offered: u64,
pub delivered: u64,
pub dropped: u64,
pub achieved: f64,
}
pub async fn drive_rated_load(
mut config: Config,
target_rate: u32,
duration: Duration,
) -> RatedLoadResult {
const WORKERS: u32 = 32;
assert!(config.siem_url.is_none(), "exactly one sink: the file sink");
let dir = tempfile::tempdir().expect("a temporary directory");
let path = dir.path().join("decisions.ndjson");
config.log_file_path = Some(path.clone());
if config.log_file_max_bytes == sample_config().log_file_max_bytes {
let records = u64::from(target_rate).saturating_mul(duration.as_secs() + 1) + 1;
let bytes = records.saturating_mul(4096).max(1 << 20);
config.log_file_max_bytes = std::num::NonZeroU64::new(bytes).unwrap();
}
probation::http::logging::set_summary_window(Duration::from_secs(3600));
let server = TestServer::start_with(config, Arc::new(SystemClock)).await;
let app = server.app(); let url = server.url("/health/live");
let offered = Arc::new(std::sync::atomic::AtomicU64::new(0));
let interval = Duration::from_secs_f64(f64::from(WORKERS) / f64::from(target_rate));
let barrier = Arc::new(tokio::sync::Barrier::new(WORKERS as usize + 1));
let workers: Vec<_> = (0..WORKERS)
.map(|_| {
let (client, url, offered) = (downstream_client(), url.clone(), Arc::clone(&offered));
let barrier = Arc::clone(&barrier);
tokio::spawn(async move {
let get = || async {
let status = client.get(&url).send().await.expect("a response").status();
assert!(status.is_success(), "/health/live answered {status}");
offered.fetch_add(1, Ordering::Relaxed);
};
get().await;
barrier.wait().await;
let started = Instant::now();
let mut next = started;
while started.elapsed() < duration {
get().await;
next += interval;
tokio::time::sleep_until(next.into()).await;
}
})
})
.collect();
barrier.wait().await;
let started = Instant::now();
for worker in workers {
worker.await.expect("a load worker");
}
let elapsed = started.elapsed();
let shutdown = tokio::spawn(server.shutdown());
let drain_limit = Instant::now() + Duration::from_secs(6);
let mut seen_summary = false;
while !seen_summary && Instant::now() < drain_limit {
tokio::time::sleep(Duration::from_millis(20)).await;
seen_summary = std::fs::read_to_string(&path)
.is_ok_and(|text| text.contains("\"event\":\"request_summary\""));
}
let dropped = app.delivery_lost_total();
assert!(
seen_summary || dropped > 0,
"the drain deadline was hit, so the reconciliation is not exact"
);
drop(app);
shutdown.await.expect("the shutdown task");
let mut rolled = path.clone().into_os_string();
rolled.push(".1");
assert!(
!Path::new(&rolled).exists(),
"the file sink rotated: the run is void, delivered would be under-counted"
);
let text = std::fs::read_to_string(&path).expect("the delivery file is readable");
let mut ids = std::collections::HashSet::new();
let mut summaries = 0;
for line in text.lines() {
let record: serde_json::Value = serde_json::from_str(line).expect("one JSON object a line");
match record["event"].as_str() {
Some("request_decided") => {
ids.insert(
record["request_id"]
.as_str()
.expect("a request_id")
.to_owned(),
);
}
Some("request_summary") => summaries += 1,
other => panic!("unexpected record kind {other:?}"),
}
}
let offered = offered.load(Ordering::Relaxed);
RatedLoadResult {
offered,
delivered: ids.len() as u64 + summaries,
dropped,
achieved: (offered - u64::from(WORKERS)) as f64 / elapsed.as_secs_f64(),
}
}