use std::convert::Infallible;
use std::task::{Context, Poll};
use std::time::Duration;
use axum::Router;
use axum::body::Body;
use axum::extract::Request;
use axum::response::Response;
use axum::routing::RouterIntoService;
use tokio::sync::{mpsc, oneshot, watch};
use tower::Service;
use crate::config::{Config, ProfileConfig};
pub struct Applied<'a> {
pub config: &'a Config,
pub profiles: &'a [ProfileConfig],
}
type Projection = fn(&Applied<'_>) -> String;
const FROZEN: &[(&str, Projection)] = &[("database.url", |a| a.config.database.url.clone())];
#[derive(Debug, thiserror::Error)]
pub enum ReloadError {
#[error(
"`{key}` cannot be changed while the server is running \
(running with `{applied}`, the file now says `{proposed}`): \
restart to apply it"
)]
Frozen {
key: String,
applied: String,
proposed: String,
},
#[error("the configuration did not load: {0}")]
Load(String),
#[error("the new configuration did not build: {0}")]
Build(String),
}
impl ReloadError {
#[must_use]
pub fn kind(&self) -> &'static str {
match self {
Self::Frozen { .. } => "frozen_key",
Self::Load(_) => "load_failed",
Self::Build(_) => "build_failed",
}
}
}
pub fn check_frozen(applied: &Applied<'_>, proposed: &Applied<'_>) -> Result<(), ReloadError> {
for (key, read) in FROZEN {
let (before, after) = (read(applied), read(proposed));
if before != after {
return Err(ReloadError::Frozen {
key: (*key).to_string(),
applied: before,
proposed: after,
});
}
}
Ok(())
}
#[derive(Debug, Clone)]
pub struct ReloadReport {
pub generation: u64,
pub profiles: Vec<String>,
pub job_kinds: Vec<&'static str>,
pub tls_reloaded: bool,
pub admin_tls_reloaded: bool,
pub listeners_rebound: Vec<&'static str>,
pub logging_reloaded: bool,
pub duration: Duration,
}
pub struct ReloadRequest {
pub respond: Option<oneshot::Sender<Result<ReloadReport, ReloadError>>>,
}
#[derive(Clone)]
pub struct ReloadHandle(mpsc::Sender<ReloadRequest>);
impl ReloadHandle {
pub fn trigger(&self) -> bool {
self.0.try_send(ReloadRequest { respond: None }).is_ok()
}
pub async fn reload(&self) -> Result<ReloadReport, ReloadError> {
let (respond, answer) = oneshot::channel();
self.0
.send(ReloadRequest {
respond: Some(respond),
})
.await
.map_err(|_| ReloadError::Load("the server is not accepting reloads".to_string()))?;
answer
.await
.map_err(|_| ReloadError::Load("the reload was abandoned".to_string()))?
}
}
pub struct Reloads(mpsc::Receiver<ReloadRequest>);
impl Reloads {
#[must_use]
pub fn none() -> Self {
Self(mpsc::channel(1).1)
}
pub(crate) async fn recv(&mut self) -> Option<ReloadRequest> {
self.0.recv().await
}
}
#[must_use]
pub fn channel() -> (ReloadHandle, Reloads) {
let (sender, receiver) = mpsc::channel(1);
(ReloadHandle(sender), Reloads(receiver))
}
#[derive(Clone)]
struct SwapService(watch::Receiver<RouterIntoService<Body>>);
impl Service<Request> for SwapService {
type Response = Response;
type Error = Infallible;
type Future = <RouterIntoService<Body> as Service<Request>>::Future;
fn poll_ready(&mut self, _context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, request: Request) -> Self::Future {
let mut current = self.0.borrow().clone();
current.call(request)
}
}
pub fn swappable(current: watch::Receiver<RouterIntoService<Body>>) -> Router {
Router::new().fallback_service(SwapService(current))
}
#[must_use]
pub fn router_channel(
initial: Router,
) -> (
watch::Sender<RouterIntoService<Body>>,
watch::Receiver<RouterIntoService<Body>>,
) {
watch::channel(initial.into_service::<Body>())
}
#[cfg(test)]
mod channel_tests {
use super::*;
#[tokio::test]
async fn a_second_trigger_while_one_is_pending_coalesces() {
let (handle, mut reloads) = channel();
assert!(handle.trigger(), "the first request is accepted");
assert!(!handle.trigger(), "the second finds one already queued");
let request = reloads.recv().await.expect("the queued request arrives");
assert!(
request.respond.is_none(),
"a signal has nobody to answer, so it asks for no channel"
);
assert!(handle.trigger());
}
#[tokio::test]
async fn a_reload_with_no_supervisor_left_fails_rather_than_hanging() {
let (handle, reloads) = channel();
drop(reloads);
let error = handle
.reload()
.await
.expect_err("there is nobody to serve the request");
assert_eq!(error.kind(), "load_failed");
assert!(
error.to_string().contains("not accepting reloads"),
"{error}"
);
}
#[tokio::test]
async fn a_source_that_never_fires_ends_at_once() {
assert!(Reloads::none().recv().await.is_none());
}
}
#[cfg(test)]
mod frozen_tests {
use super::*;
use crate::config::ProfileSections;
fn profile(name: &str) -> ProfileConfig {
ProfileConfig {
name: name.to_string(),
sections: ProfileSections::default(),
}
}
fn refuse(
mutate: impl FnOnce(&mut Config, &mut Vec<ProfileConfig>),
) -> Result<(), ReloadError> {
let applied = Config::default();
let applied_profiles = vec![profile("le")];
let mut proposed = Config::default();
let mut proposed_profiles = vec![profile("le")];
mutate(&mut proposed, &mut proposed_profiles);
check_frozen(
&Applied {
config: &applied,
profiles: &applied_profiles,
},
&Applied {
config: &proposed,
profiles: &proposed_profiles,
},
)
}
fn refused_key(result: Result<(), ReloadError>) -> String {
match result {
Err(ReloadError::Frozen { key, .. }) => key,
Err(other) => panic!("expected a frozen-key refusal, got {other}"),
Ok(()) => panic!("expected a refusal, the change was allowed"),
}
}
#[test]
fn every_frozen_key_is_refused_by_its_own_name() {
#[allow(clippy::type_complexity)]
let cases: Vec<(&str, Box<dyn Fn(&mut Config, &mut Vec<ProfileConfig>)>)> = vec![(
"database.url",
Box::new(|c: &mut Config, _: &mut Vec<ProfileConfig>| {
c.database.url = "sqlite://other.db".to_string();
}),
)];
for (key, mutate) in &cases {
let refused = refused_key(refuse(|config, profiles| mutate(config, profiles)));
assert_eq!(
&refused.as_str(),
key,
"changing `{key}` must be refused naming `{key}`, not `{refused}`",
);
}
let covered: std::collections::BTreeSet<&str> = cases.iter().map(|(key, _)| *key).collect();
let table: std::collections::BTreeSet<&str> = FROZEN.iter().map(|(key, _)| *key).collect();
assert_eq!(
covered, table,
"every FROZEN entry needs a case here, and every case needs an entry",
);
}
#[test]
fn an_unchanged_configuration_is_allowed() {
assert!(refuse(|_, _| {}).is_ok());
}
#[test]
fn the_profile_set_its_signers_and_the_egress_all_reload() {
assert!(refuse(|_, profiles| profiles.push(profile("staging"))).is_ok());
assert!(refuse(|_, profiles| profiles[0].name = "staging".to_string()).is_ok());
assert!(refuse(|_, profiles| profiles.clear()).is_ok());
assert!(
refuse(|_, profiles| profiles[0].sections.signer.backend = "custom".to_string())
.is_ok()
);
assert!(
refuse(|_, profiles| profiles[0].sections.signer.local_ca.leaf_validity_days = 30)
.is_ok()
);
assert!(refuse(|c, _| c.dns.resolver = Some("192.0.2.1:53".to_string())).is_ok());
assert!(refuse(|c, _| c.proxy.https_url = "http://proxy.example:3128".to_string()).is_ok());
}
#[test]
fn nothing_about_the_signer_sections_is_refused_any_more() {
assert!(refuse(|config, _| config.signer.backend = "custom".to_string()).is_ok());
assert!(
refuse(|config, profiles| {
config.signer.backend = "custom".to_string();
profiles[0].sections.signer.backend = "custom".to_string();
})
.is_ok()
);
}
#[test]
fn every_jobs_key_is_reloadable() {
for mutate in [
|c: &mut Config| c.jobs.poll_interval_ms += 1,
|c: &mut Config| c.jobs.max_concurrent += 1,
|c: &mut Config| c.jobs.max_attempts += 1,
|c: &mut Config| c.jobs.retry_base_seconds += 1,
|c: &mut Config| c.jobs.retry_max_seconds += 1,
|c: &mut Config| c.jobs.lease_seconds += 1,
|c: &mut Config| c.jobs.retention_days += 1,
] {
assert!(
refuse(|config, _| mutate(config)).is_ok(),
"nothing in `[jobs]` is snapshotted at spawn any more, so no key \
in it is frozen",
);
}
}
#[test]
fn every_logging_key_is_reloadable() {
for mutate in [
|c: &mut Config| c.logging.filter = "acme_proxy=debug".to_string(),
|c: &mut Config| c.logging.json_format = true,
|c: &mut Config| c.logging.flatten_event = true,
|c: &mut Config| c.logging.target = "stderr".to_string(),
|c: &mut Config| c.logging.ansi = false,
|c: &mut Config| c.logging.span_events = "close".to_string(),
] {
assert!(
refuse(|config, _| mutate(config)).is_ok(),
"the layer stack is swapped whole, so no `[logging]` key is frozen",
);
}
}
#[test]
fn every_listener_key_is_reloadable() {
for mutate in [
|c: &mut Config| c.server.bind_address = "127.0.0.1:9999".to_string(),
|c: &mut Config| c.server.tls.enabled = !c.server.tls.enabled,
|c: &mut Config| c.admin.enabled = !c.admin.enabled,
|c: &mut Config| c.admin.bind_address = "127.0.0.1:9998".to_string(),
|c: &mut Config| c.admin.tls.enabled = !c.admin.tls.enabled,
|c: &mut Config| c.metrics.enabled = !c.metrics.enabled,
|c: &mut Config| c.metrics.bind_address = "127.0.0.1:9997".to_string(),
] {
assert!(
refuse(|config, _| mutate(config)).is_ok(),
"a socket is replaceable now, so no bind address or listener \
switch is frozen",
);
}
}
#[test]
fn a_refusal_names_the_key_and_both_values() {
let error = refuse(|config, _| {
config.database.url = "sqlite://elsewhere.db".to_string();
})
.expect_err("a changed database URL is refused");
let rendered = error.to_string();
assert!(rendered.contains("database.url"), "{rendered}");
assert!(rendered.contains("sqlite://sqlite.db"), "{rendered}");
assert!(rendered.contains("sqlite://elsewhere.db"), "{rendered}");
assert_eq!(error.kind(), "frozen_key");
}
#[test]
fn the_other_failures_describe_themselves() {
let load = ReloadError::Load("no such file".to_string());
assert_eq!(load.kind(), "load_failed");
assert!(load.to_string().contains("did not load"), "{load}");
let build = ReloadError::Build("bad filter rule".to_string());
assert_eq!(build.kind(), "build_failed");
assert!(build.to_string().contains("did not build"), "{build}");
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::to_bytes;
use axum::routing::get;
use http_body_util::BodyExt;
use tower::ServiceExt;
fn answering(body: &'static str) -> Router {
Router::new().route("/", get(move || async move { body }))
}
async fn body_of(response: Response) -> String {
String::from_utf8(
to_bytes(response.into_body(), usize::MAX)
.await
.unwrap()
.to_vec(),
)
.unwrap()
}
#[tokio::test]
async fn a_request_after_a_swap_reaches_the_new_router() {
let (sender, receiver) = router_channel(answering("first"));
let app = swappable(receiver);
let response = app
.clone()
.oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(body_of(response).await, "first");
sender.send_replace(answering("second").into_service::<Body>());
let response = app
.oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(body_of(response).await, "second");
}
#[tokio::test]
async fn the_wrapper_never_answers_in_place_of_the_router_it_holds() {
let inner = answering("routed").fallback(|| async { "inner fallback" });
let (_sender, receiver) = router_channel(inner);
let response = swappable(receiver)
.oneshot(
Request::builder()
.uri("/nothing-here")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(body_of(response).await, "inner fallback");
}
#[tokio::test]
async fn a_head_request_keeps_its_content_length_and_loses_its_body() {
let (_sender, receiver) = router_channel(answering("first"));
let response = swappable(receiver)
.oneshot(
Request::builder()
.method("HEAD")
.uri("/")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(
response.headers().get("content-length").unwrap(),
"5",
"the length of the body a GET would have returned"
);
assert!(
response
.into_body()
.collect()
.await
.unwrap()
.to_bytes()
.is_empty(),
"a HEAD response carries no body"
);
}
}