mod signature;
pub use signature::{SIGNATURE_HEADER, ServiceSigner, SignedRequests};
use axum::{
body::Body,
extract::{FromRequestParts, Request},
http::{
HeaderMap, StatusCode,
header::{AUTHORIZATION, HeaderName, HeaderValue, WWW_AUTHENTICATE},
request::Parts,
},
response::{IntoResponse, Response},
};
use sha2::{Digest, Sha256};
use signature::Verifier;
use std::{
future::Future,
path::{Path, PathBuf},
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use subtle::{Choice, ConstantTimeEq};
use tower::{Layer, Service};
use zeroize::Zeroizing;
pub const MIN_SECRET_LEN: usize = 32;
#[derive(Debug, thiserror::Error)]
pub enum ServiceSecretError {
#[error("cannot read the secret for `{name}` from {path}: {source}")]
Read {
name: String,
path: PathBuf,
#[source]
source: std::io::Error,
},
#[error("the secret for `{name}` is {len} bytes; at least {MIN_SECRET_LEN} are required")]
TooShort {
name: String,
len: usize,
},
}
#[derive(Clone)]
pub struct ServiceSecret {
name: Arc<str>,
digest: [u8; 32],
signing_key: Zeroizing<[u8; 32]>,
}
impl std::fmt::Debug for ServiceSecret {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ServiceSecret")
.field("name", &self.name)
.finish_non_exhaustive()
}
}
impl ServiceSecret {
pub fn new(name: impl Into<Arc<str>>, secret: &str) -> Result<Self, ServiceSecretError> {
let name = name.into();
let secret = secret.trim();
if secret.len() < MIN_SECRET_LEN {
return Err(ServiceSecretError::TooShort {
name: name.to_string(),
len: secret.len(),
});
}
Ok(Self {
name,
digest: Sha256::digest(secret.as_bytes()).into(),
signing_key: signature::derive_key(secret.as_bytes()),
})
}
pub fn from_file(
name: impl Into<Arc<str>>,
path: impl AsRef<Path>,
) -> Result<Self, ServiceSecretError> {
let name = name.into();
let path = path.as_ref();
let secret = Zeroizing::new(std::fs::read_to_string(path).map_err(|source| {
ServiceSecretError::Read {
name: name.to_string(),
path: path.to_path_buf(),
source,
}
})?);
Self::new(name, &secret)
}
pub fn from_env(
name: impl Into<Arc<str>>,
var: &str,
) -> Result<Option<Self>, ServiceSecretError> {
let name = name.into();
if let Some(path) = std::env::var_os(format!("{var}_FILE")) {
return Self::from_file(name, path).map(Some);
}
match std::env::var(var) {
Ok(secret) => Self::new(name, &Zeroizing::new(secret)).map(Some),
Err(_) => Ok(None),
}
}
pub fn name(&self) -> &str {
&self.name
}
pub fn signer(&self) -> ServiceSigner {
ServiceSigner::from_key(self.signing_key.clone())
}
}
#[derive(Clone, Debug, Default)]
pub struct ServiceSecrets {
secrets: Vec<ServiceSecret>,
}
impl ServiceSecrets {
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with(mut self, secret: ServiceSecret) -> Self {
self.secrets.push(secret);
self
}
pub fn is_configured(&self) -> bool {
!self.secrets.is_empty()
}
pub fn identify(&self, presented: &[u8]) -> Option<ServiceCaller> {
let presented: [u8; 32] = Sha256::digest(presented).into();
self.first_match(|secret| secret.digest.ct_eq(&presented))
}
fn first_match(&self, hit: impl Fn(&ServiceSecret) -> Choice) -> Option<ServiceCaller> {
let mut found: Option<&ServiceSecret> = None;
for secret in &self.secrets {
if bool::from(hit(secret)) && found.is_none() {
found = Some(secret);
}
}
found.map(|secret| ServiceCaller(Arc::clone(&secret.name)))
}
}
impl FromIterator<ServiceSecret> for ServiceSecrets {
fn from_iter<I: IntoIterator<Item = ServiceSecret>>(iter: I) -> Self {
Self {
secrets: iter.into_iter().collect(),
}
}
}
impl Extend<ServiceSecret> for ServiceSecrets {
fn extend<I: IntoIterator<Item = ServiceSecret>>(&mut self, iter: I) {
self.secrets.extend(iter);
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ServiceCaller(Arc<str>);
impl ServiceCaller {
pub fn name(&self) -> &str {
&self.0
}
}
impl<S: Send + Sync> FromRequestParts<S> for ServiceCaller {
type Rejection = StatusCode;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<ServiceCaller>()
.cloned()
.ok_or(StatusCode::INTERNAL_SERVER_ERROR)
}
}
#[derive(Clone, Debug)]
enum Mode {
Bearer,
Header(HeaderName),
Signed(Arc<Verifier>),
}
impl Mode {
fn presented<'h>(&self, headers: &'h HeaderMap) -> Option<&'h [u8]> {
match self {
Self::Bearer => {
let value = headers.get(AUTHORIZATION)?.as_bytes();
let (scheme, rest) = value.split_at_checked(7)?;
scheme
.eq_ignore_ascii_case(b"bearer ")
.then(|| rest.trim_ascii())
}
Self::Header(name) => Some(headers.get(name)?.as_bytes().trim_ascii()),
Self::Signed(_) => None,
}
}
}
#[derive(Clone, Debug)]
pub struct ServiceSecretLayer {
secrets: Arc<ServiceSecrets>,
mode: Mode,
}
impl ServiceSecretLayer {
pub fn bearer(secrets: ServiceSecrets) -> Self {
Self::with_mode(secrets, Mode::Bearer)
}
pub fn header(secrets: ServiceSecrets, name: HeaderName) -> Self {
Self::with_mode(secrets, Mode::Header(name))
}
pub fn signed(secrets: ServiceSecrets, config: SignedRequests) -> Self {
Self::with_mode(secrets, Mode::Signed(Arc::new(Verifier::new(config))))
}
fn with_mode(secrets: ServiceSecrets, mode: Mode) -> Self {
Self {
secrets: Arc::new(secrets),
mode,
}
}
}
impl<S> Layer<S> for ServiceSecretLayer {
type Service = ServiceSecretService<S>;
fn layer(&self, inner: S) -> Self::Service {
ServiceSecretService {
inner,
secrets: Arc::clone(&self.secrets),
mode: self.mode.clone(),
}
}
}
#[derive(Clone, Debug)]
pub struct ServiceSecretService<S> {
inner: S,
secrets: Arc<ServiceSecrets>,
mode: Mode,
}
fn refusal(secrets: &ServiceSecrets, mode: &Mode) -> Response {
if !secrets.is_configured() {
return StatusCode::NOT_FOUND.into_response();
}
let mut response = StatusCode::UNAUTHORIZED.into_response();
if matches!(mode, Mode::Bearer) {
response
.headers_mut()
.insert(WWW_AUTHENTICATE, HeaderValue::from_static("Bearer"));
}
response
}
impl<S> Service<Request<Body>> for ServiceSecretService<S>
where
S: Service<Request<Body>, Response = Response> + Send + Clone + 'static,
S::Future: Send + 'static,
{
type Response = Response;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Response, S::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: Request<Body>) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
let secrets = Arc::clone(&self.secrets);
if let Mode::Signed(verifier) = &self.mode {
let verifier = Arc::clone(verifier);
let mode = self.mode.clone();
return Box::pin(async move {
if !secrets.is_configured() {
return Ok(refusal(&secrets, &mode));
}
match verifier.verify(&secrets, req).await {
Ok(req) => inner.call(req).await,
Err(status) if status == StatusCode::UNAUTHORIZED => {
Ok(refusal(&secrets, &mode))
}
Err(status) => Ok(status.into_response()),
}
});
}
let caller = self
.mode
.presented(req.headers())
.and_then(|presented| secrets.identify(presented));
let Some(caller) = caller else {
let refused = refusal(&secrets, &self.mode);
return Box::pin(async move { Ok(refused) });
};
let mut req = req;
req.extensions_mut().insert(caller);
Box::pin(inner.call(req))
}
}
#[cfg(test)]
mod tests;