#![warn(missing_docs)]
use std::task::{Context, Poll};
use tonic::{
Status,
metadata::{Ascii, MetadataValue},
};
use tower::{Layer, Service};
use super::config::{API_TOKEN_METADATA_KEY, APP_API_TOKEN_ENV};
#[derive(Clone, Debug, Default)]
pub struct ApiTokenInterceptor {
token: Option<MetadataValue<Ascii>>,
}
impl ApiTokenInterceptor {
pub fn new(token: Option<String>) -> Self {
Self::try_new(token).unwrap_or_default()
}
pub fn try_new(token: Option<String>) -> Result<Self, crate::error::Error> {
let token = match token {
Some(t) if !t.is_empty() => Some(t.parse()?),
_ => None,
};
Ok(Self { token })
}
pub fn has_token(&self) -> bool {
self.token.is_some()
}
}
impl tonic::service::Interceptor for ApiTokenInterceptor {
fn call(&mut self, mut request: tonic::Request<()>) -> Result<tonic::Request<()>, Status> {
if let Some(token) = &self.token {
request
.metadata_mut()
.insert(API_TOKEN_METADATA_KEY, token.clone());
}
Ok(request)
}
}
#[derive(Clone, Debug, Default)]
pub struct AppApiTokenLayer {
expected: Option<String>,
}
impl AppApiTokenLayer {
pub fn new(expected: Option<String>) -> Self {
let expected = match expected {
Some(s) if !s.is_empty() => Some(s),
_ => None,
};
Self { expected }
}
pub fn from_env() -> Self {
Self::new(std::env::var(APP_API_TOKEN_ENV).ok())
}
pub fn is_enforcing(&self) -> bool {
self.expected.is_some()
}
}
impl<S> Layer<S> for AppApiTokenLayer {
type Service = AppApiTokenService<S>;
fn layer(&self, inner: S) -> Self::Service {
AppApiTokenService {
inner,
expected: self.expected.clone(),
}
}
}
#[derive(Clone, Debug)]
pub struct AppApiTokenService<S> {
inner: S,
expected: Option<String>,
}
impl<S, ReqBody, ResBody> Service<http::Request<ReqBody>> for AppApiTokenService<S>
where
S: Service<http::Request<ReqBody>, Response = http::Response<ResBody>> + Clone + Send + 'static,
S::Future: Send + 'static,
ReqBody: Send + 'static,
ResBody: Default,
{
type Response = S::Response;
type Error = S::Error;
type Future = futures::future::BoxFuture<'static, Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: http::Request<ReqBody>) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
let expected = self.expected.clone();
Box::pin(async move {
if let Some(expected_token) = expected {
let presented = req
.headers()
.get(API_TOKEN_METADATA_KEY)
.and_then(|v| v.to_str().ok());
if presented != Some(expected_token.as_str()) {
let response = http::Response::builder()
.status(http::StatusCode::UNAUTHORIZED)
.header(
"grpc-status",
(tonic::Code::Unauthenticated as i32).to_string(),
)
.header("grpc-message", "invalid or missing dapr-api-token")
.body(ResBody::default())
.expect("static response is valid");
return Ok(response);
}
}
inner.call(req).await
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use tonic::Request;
use tonic::service::Interceptor;
#[test]
fn interceptor_no_token_is_noop() {
let mut interceptor = ApiTokenInterceptor::new(None);
let req = interceptor.call(Request::new(())).unwrap();
assert!(req.metadata().get(API_TOKEN_METADATA_KEY).is_none());
}
#[test]
fn interceptor_empty_token_is_noop() {
let mut interceptor = ApiTokenInterceptor::new(Some(String::new()));
assert!(!interceptor.has_token());
let req = interceptor.call(Request::new(())).unwrap();
assert!(req.metadata().get(API_TOKEN_METADATA_KEY).is_none());
}
#[test]
fn interceptor_injects_token() {
let mut interceptor = ApiTokenInterceptor::new(Some("abc".to_string()));
let req = interceptor.call(Request::new(())).unwrap();
assert_eq!(req.metadata().get(API_TOKEN_METADATA_KEY).unwrap(), "abc");
}
#[test]
fn interceptor_rejects_invalid_metadata() {
assert!(matches!(
ApiTokenInterceptor::try_new(Some("bad\nvalue".to_string())),
Err(crate::error::Error::InvalidMetadata)
));
}
#[test]
fn app_layer_permissive_when_no_token() {
let layer = AppApiTokenLayer::new(None);
assert!(!layer.is_enforcing());
}
#[test]
fn app_layer_enforces_when_token_set() {
let layer = AppApiTokenLayer::new(Some("token".to_string()));
assert!(layer.is_enforcing());
}
#[tokio::test]
async fn app_layer_rejects_unauthenticated_axum_requests() {
use axum::{Router, routing::get};
use http_body_util::BodyExt;
use tower::ServiceExt;
let app: Router = Router::new()
.route("/secret", get(|| async { "ok" }))
.layer(AppApiTokenLayer::new(Some("expected".to_string())));
let resp = app
.clone()
.oneshot(
http::Request::builder()
.uri("/secret")
.body(axum::body::Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), http::StatusCode::UNAUTHORIZED);
let resp = app
.clone()
.oneshot(
http::Request::builder()
.uri("/secret")
.header(API_TOKEN_METADATA_KEY, "wrong")
.body(axum::body::Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), http::StatusCode::UNAUTHORIZED);
let resp = app
.oneshot(
http::Request::builder()
.uri("/secret")
.header(API_TOKEN_METADATA_KEY, "expected")
.body(axum::body::Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), http::StatusCode::OK);
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(&body[..], b"ok");
}
#[tokio::test]
async fn app_layer_permissive_passes_through() {
use axum::{Router, routing::get};
use tower::ServiceExt;
let app: Router = Router::new()
.route("/open", get(|| async { "ok" }))
.layer(AppApiTokenLayer::new(None));
let resp = app
.oneshot(
http::Request::builder()
.uri("/open")
.body(axum::body::Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), http::StatusCode::OK);
}
}