use std::sync::Arc;
use axum::{
body::Body,
extract::State,
http::{HeaderName, HeaderValue, Request, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
Json,
};
use super::request_id::{RequestId, REQUEST_ID_HEADER};
pub const AUTH_HEADER: HeaderName = HeaderName::from_static("authorization");
#[derive(Clone, Default)]
pub struct AuthConfig {
pub expected: Arc<Option<String>>,
expected_digest: Option<(Arc<str>, [u8; 32])>,
}
impl AuthConfig {
pub fn new(expected: Option<String>) -> Self {
let expected_digest = expected
.as_deref()
.map(|token| (Arc::from(token), digest_of(token)));
Self {
expected: Arc::new(expected),
expected_digest,
}
}
pub fn resolve_api_key_with<F>(get_env: F) -> Option<String>
where
F: Fn(&str) -> Result<String, std::env::VarError>,
{
get_env("OPENKIND_API_KEY")
.ok()
.filter(|s| !s.is_empty())
.or_else(|| {
get_env("OPENDECISION_API_KEY")
.ok()
.filter(|s| !s.is_empty())
})
.or_else(|| get_env("TYPESAFE_API_KEY").ok().filter(|s| !s.is_empty()))
.or_else(|| get_env("OPENPICK_API_KEY").ok().filter(|s| !s.is_empty()))
}
pub fn from_env() -> Self {
Self::new(Self::resolve_api_key_with(|k| std::env::var(k)))
}
pub fn is_required(&self) -> bool {
self.expected.is_some()
}
pub(crate) fn token_matches(&self, supplied: &str) -> bool {
use subtle::ConstantTimeEq;
let Some(expected) = self.expected.as_deref() else {
return false;
};
let expected_digest = self
.expected_digest
.as_ref()
.filter(|(cached, _)| cached.as_ref() == expected)
.map(|(_, digest)| *digest)
.unwrap_or_else(|| digest_of(expected));
digest_of(supplied).ct_eq(&expected_digest).into()
}
}
pub async fn auth_layer(
State(auth): State<AuthConfig>,
req: Request<Body>,
next: Next,
) -> Response {
let path = req.uri().path();
if !auth.is_required()
|| path == "/health"
|| path == "/metrics"
|| path == "/playground"
|| req.method() == axum::http::Method::OPTIONS
{
return next.run(req).await;
}
let supplied = req
.headers()
.get(&AUTH_HEADER)
.and_then(|v| v.to_str().ok())
.and_then(|s| {
let (scheme, token) = s.split_once(' ')?;
scheme.eq_ignore_ascii_case("Bearer").then_some(token)
});
let ok = supplied.is_some_and(|token| auth.token_matches(token));
if !ok {
let body = Json(serde_json::json!({
"error": {
"code": "unauthorized",
"message": "missing or invalid API key",
}
}));
let mut resp = (StatusCode::UNAUTHORIZED, body).into_response();
if let Ok(v) = HeaderValue::from_str("Bearer") {
resp.headers_mut()
.insert(axum::http::header::WWW_AUTHENTICATE, v);
}
if let Some(req_id) = req.extensions().get::<RequestId>() {
if let Ok(v) = HeaderValue::from_str(&req_id.0) {
resp.headers_mut().insert(REQUEST_ID_HEADER.clone(), v);
}
}
return resp;
}
next.run(req).await
}
fn digest_of(token: &str) -> [u8; 32] {
let digest = ring::digest::digest(&ring::digest::SHA256, token.as_bytes());
let mut out = [0u8; 32];
out.copy_from_slice(digest.as_ref());
out
}
pub fn secure_token_eq(a: &str, b: &str) -> bool {
use subtle::ConstantTimeEq;
digest_of(a).ct_eq(&digest_of(b)).into()
}
async fn auth_layer_dummy_handler() -> StatusCode {
StatusCode::OK
}
pub fn auth_layer_for(auth: AuthConfig) -> axum::Router {
axum::Router::new()
.route("/", axum::routing::get(auth_layer_dummy_handler))
.layer(axum::middleware::from_fn_with_state(auth, auth_layer))
}