use std::path::{Path, PathBuf};
use axum::extract::{Request, State};
use axum::http::HeaderValue;
use axum::http::header::{COOKIE, SET_COOKIE};
use axum::middleware::Next;
use axum::response::{IntoResponse, Redirect, Response};
use serde::{Deserialize, Serialize};
use crate::crypto::constant_time_eq;
use crate::queue::unix_now;
use crate::{AppState, Error, Result};
const BYPASS_COOKIE: &str = "renox_maintenance";
#[derive(Debug, Clone, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Down {
#[serde(with = "chrono::serde::ts_seconds")]
pub since: crate::db::DateTime,
pub retry: Option<u64>,
pub secret: Option<String>,
}
pub(crate) fn file(storage: &Path) -> PathBuf {
storage.join("framework").join("down")
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct DownOptions {
secret: Option<String>,
retry: Option<u64>,
}
impl DownOptions {
pub fn new() -> Self {
Self::default()
}
pub fn secret(mut self, secret: impl Into<String>) -> Self {
self.secret = Some(secret.into());
self
}
pub fn retry(mut self, seconds: u64) -> Self {
self.retry = Some(seconds);
self
}
}
pub fn down(storage: &Path, options: DownOptions) -> Result {
let path = file(storage);
if let Some(dir) = path.parent() {
std::fs::create_dir_all(dir)?;
}
let state = Down {
since: crate::db::from_unix(unix_now()),
retry: options.retry,
secret: options.secret,
};
std::fs::write(path, serde_json::to_string(&state)?)?;
Ok(())
}
pub fn up(storage: &Path) -> Result<bool> {
match std::fs::remove_file(file(storage)) {
Ok(()) => Ok(true),
Err(err) if err.kind() == std::io::ErrorKind::NotFound => Ok(false),
Err(err) => Err(err.into()),
}
}
pub fn status(storage: &Path) -> Option<Down> {
let text = std::fs::read_to_string(file(storage)).ok()?;
serde_json::from_str(&text).ok()
}
fn bypass_token(state: &AppState, secret: &str) -> String {
crate::signed::signature(state, &format!("maintenance-bypass:{secret}"))
}
fn has_bypass(req: &Request, token: &str) -> bool {
req.headers()
.get_all(COOKIE)
.iter()
.filter_map(|v| v.to_str().ok())
.flat_map(|v| v.split(';'))
.filter_map(|pair| pair.trim().split_once('='))
.any(|(name, value)| name == BYPASS_COOKIE && constant_time_eq(value, token))
}
pub(crate) async fn middleware(
State(state): State<AppState>,
req: Request,
next: Next,
) -> Response {
let Some(down) = status(&state.config.storage_path) else {
return next.run(req).await;
};
if state
.security
.is_webhook(req.extensions().get::<axum::extract::MatchedPath>())
{
return next.run(req).await;
}
if let Some(secret) = &down.secret {
if has_bypass(&req, &bypass_token(&state, secret)) {
return next.run(req).await;
}
if req.uri().path().trim_start_matches('/') == secret {
let mut res = Redirect::to("/").into_response();
let secure = if state.config.url.starts_with("https://") {
"; Secure"
} else {
""
};
let token = bypass_token(&state, secret);
let cookie = format!("{BYPASS_COOKIE}={token}; Path=/; HttpOnly; SameSite=Lax{secure}");
if let Ok(value) = HeaderValue::from_str(&cookie) {
res.headers_mut().append(SET_COOKIE, value);
}
return res;
}
}
let mut res = Error::ServiceUnavailable.into_response();
if let Some(retry) = down.retry {
res.headers_mut().insert("retry-after", retry.into());
}
res
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn up_reports_a_marker_it_cant_remove() {
let dir = tempfile::tempdir().unwrap();
assert!(!up(dir.path()).unwrap(), "nothing to bring up");
std::fs::create_dir_all(file(dir.path())).unwrap();
assert!(up(dir.path()).is_err());
}
}