use crate::error::ApiErrorResponse;
use crate::router::{
MATCHED_KB, MATCHED_KB_EVENTS, MATCHED_KB_INDEX, MATCHED_KB_OBJECT, MATCHED_KB_SEARCH,
MATCHED_KBS,
};
use crate::state::AppState;
use axum::body::Body;
use axum::extract::{MatchedPath, State};
use axum::http::{Method, Request, StatusCode, header::WWW_AUTHENTICATE};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use notedthat_core::{Principal, Schemes};
use tower_http::request_id::RequestId;
const ANONYMOUS_REACHABLE: &[(&Method, &str)] = &[
(&Method::GET, MATCHED_KBS),
(&Method::HEAD, MATCHED_KBS),
(&Method::GET, MATCHED_KB),
(&Method::HEAD, MATCHED_KB),
(&Method::GET, MATCHED_KB_OBJECT),
(&Method::HEAD, MATCHED_KB_OBJECT),
(&Method::POST, MATCHED_KB_SEARCH),
(&Method::GET, MATCHED_KB_EVENTS),
(&Method::GET, MATCHED_KB_INDEX),
];
pub async fn auth_middleware(
State(state): State<AppState>,
mut req: Request<Body>,
next: Next,
) -> Response {
let request_id = extract_request_id(&req);
let response = match state
.authenticator
.resolve(req.headers(), Schemes::Bearer)
.await
{
Err(CredentialRefused) => ApiErrorResponse::unauthorized(request_id).into_response(),
Ok(principal) if principal.is_anonymous() && !anonymous_may_reach(&req) => {
ApiErrorResponse::unauthorized(request_id).into_response()
}
Ok(principal) => {
req.extensions_mut().insert(principal);
next.run(req).await
}
};
with_bearer_challenge(&state, response)
}
pub(crate) fn with_bearer_challenge(state: &AppState, mut response: Response) -> Response {
if response.status() == StatusCode::UNAUTHORIZED
&& let Some(challenge) = state.authenticator.bearer_challenge()
{
response.headers_mut().insert(WWW_AUTHENTICATE, challenge);
}
response
}
fn anonymous_may_reach<B>(req: &Request<B>) -> bool {
let Some(matched) = req.extensions().get::<MatchedPath>() else {
return false;
};
let matched = matched.as_str();
ANONYMOUS_REACHABLE
.iter()
.any(|(method, route)| *method == req.method() && *route == matched)
}
pub use notedthat_core::CredentialRefused;
pub fn principal<B>(req: &Request<B>) -> Principal {
req.extensions()
.get::<Principal>()
.cloned()
.unwrap_or(Principal::Anyone)
}
pub use notedthat_core::is_internal_path;
pub fn extract_request_id<B>(req: &Request<B>) -> String {
req.extensions()
.get::<RequestId>()
.and_then(|r| r.header_value().to_str().ok())
.map_or_else(
|| {
tracing::warn!("request_id missing from Extensions — generating fallback");
uuid::Uuid::now_v7().to_string()
},
str::to_string,
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::testing::InMemoryStorage;
use axum::middleware::from_fn_with_state;
use axum::response::IntoResponse;
use axum::routing::get;
use axum::{Router, body::Body, http::StatusCode};
use std::collections::BTreeMap;
use std::sync::Arc;
use tower::util::ServiceExt;
fn test_state(token: &str) -> AppState {
let (indexer_tx, _rx) = tokio::sync::mpsc::channel(1024);
AppState {
storage: Arc::new(InMemoryStorage::default()),
declared_kbs: Arc::new(BTreeMap::new()),
access_policies: Arc::new(BTreeMap::new()),
kb_details: Arc::new(BTreeMap::new()),
authenticator: Arc::new(notedthat_core::Authenticator::new(token)),
max_body_size: 16 * 1024 * 1024,
max_patchable_size: 16 * 1024 * 1024,
indexer_tx,
searcher: Arc::new(crate::testing::NoopSearcher),
events: None,
index_health: Arc::new(notedthat_indexer::IndexHealth::new()),
}
}
fn app(token: &str) -> Router {
let state = test_state(token);
Router::new()
.route("/protected", get(|| async { "secret".into_response() }))
.layer(from_fn_with_state(state.clone(), auth_middleware))
.with_state(state)
}
#[tokio::test]
async fn no_path_is_exempt_once_the_request_reaches_this_layer() {
for uri in ["/healthz", "/readyz", "/llms.txt", "/protected"] {
let resp = app("my-token")
.oneshot(Request::builder().uri(uri).body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(
resp.status(),
StatusCode::UNAUTHORIZED,
"{uri} must not be exempted by the auth layer itself"
);
}
}
#[tokio::test]
async fn test_rejects_missing_auth() {
let resp = app("my-token")
.oneshot(
Request::builder()
.uri("/protected")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn test_rejects_wrong_token() {
let resp = app("real-token")
.oneshot(
Request::builder()
.uri("/protected")
.header("authorization", "Bearer wrong-token")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn test_accepts_correct_token() {
let resp = app("my-token")
.oneshot(
Request::builder()
.uri("/protected")
.header("authorization", "Bearer my-token")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_accepts_lowercase_bearer_scheme() {
let resp = app("my-token")
.oneshot(
Request::builder()
.uri("/protected")
.header("authorization", "bearer my-token")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
}