Skip to main content

notedthat_api_http/
middleware.rs

1//! Static-Bearer authentication middleware for the `NotedThat` API.
2
3use crate::error::ApiErrorResponse;
4use crate::state::AppState;
5use axum::body::Body;
6use axum::extract::State;
7use axum::http::Request;
8use axum::middleware::Next;
9use axum::response::Response;
10use notedthat_core::{extract_bearer_from_header, verify_bearer_token};
11use tower_http::request_id::RequestId;
12
13/// Health-check paths that bypass Bearer authentication.
14const AUTH_EXEMPT_PATHS: &[&str] = &["/healthz", "/readyz"];
15
16/// Axum middleware that validates the `Authorization: Bearer <token>` header.
17///
18/// Requests to health-check paths pass through without authentication.
19/// All other requests must present a valid Bearer token that matches
20/// `state.bearer_token` (compared in constant time).
21pub async fn auth_middleware(
22    State(state): State<AppState>,
23    req: Request<Body>,
24    next: Next,
25) -> Result<Response, ApiErrorResponse> {
26    let path = req.uri().path();
27    if AUTH_EXEMPT_PATHS.contains(&path) {
28        return Ok(next.run(req).await);
29    }
30
31    let request_id = extract_request_id(&req);
32
33    let header_value = req
34        .headers()
35        .get("authorization")
36        .and_then(|v| v.to_str().ok());
37    let Some(token) = header_value.and_then(extract_bearer_from_header) else {
38        return Err(ApiErrorResponse::unauthorized(request_id));
39    };
40
41    if !verify_bearer_token(token, &state.bearer_token) {
42        return Err(ApiErrorResponse::unauthorized(request_id));
43    }
44
45    Ok(next.run(req).await)
46}
47
48/// Extract the `x-request-id` value from request extensions, falling back to a
49/// generated UUID if the `SetRequestId` middleware hasn't run yet.
50pub fn extract_request_id<B>(req: &Request<B>) -> String {
51    req.extensions()
52        .get::<RequestId>()
53        .and_then(|r| r.header_value().to_str().ok())
54        .map_or_else(
55            || {
56                tracing::warn!("request_id missing from Extensions — generating fallback");
57                uuid::Uuid::now_v7().to_string()
58            },
59            str::to_string,
60        )
61}
62
63#[cfg(test)]
64mod tests {
65    use super::*;
66    use crate::testing::InMemoryStorage;
67    use axum::middleware::from_fn_with_state;
68    use axum::response::IntoResponse;
69    use axum::routing::get;
70    use axum::{Router, body::Body, http::StatusCode};
71    use std::collections::BTreeMap;
72    use std::sync::Arc;
73    use tower::util::ServiceExt;
74
75    fn test_state(token: &str) -> AppState {
76        let (indexer_tx, _rx) = tokio::sync::mpsc::channel(1024);
77        AppState {
78            storage: Arc::new(InMemoryStorage::default()),
79            declared_kbs: Arc::new(BTreeMap::new()),
80            bearer_token: Arc::new(token.to_string()),
81            max_body_size: 16 * 1024 * 1024,
82            max_patchable_size: 16 * 1024 * 1024,
83            indexer_tx,
84            searcher: Arc::new(crate::testing::NoopSearcher),
85        }
86    }
87
88    fn app(token: &str) -> Router {
89        let state = test_state(token);
90        Router::new()
91            .route("/healthz", get(|| async { "ok" }))
92            .route("/protected", get(|| async { "secret".into_response() }))
93            .layer(from_fn_with_state(state.clone(), auth_middleware))
94            .with_state(state)
95    }
96
97    #[tokio::test]
98    async fn test_healthz_bypasses_auth() {
99        let resp = app("my-token")
100            .oneshot(
101                Request::builder()
102                    .uri("/healthz")
103                    .body(Body::empty())
104                    .unwrap(),
105            )
106            .await
107            .unwrap();
108        assert_eq!(resp.status(), StatusCode::OK);
109    }
110
111    #[tokio::test]
112    async fn test_rejects_missing_auth() {
113        let resp = app("my-token")
114            .oneshot(
115                Request::builder()
116                    .uri("/protected")
117                    .body(Body::empty())
118                    .unwrap(),
119            )
120            .await
121            .unwrap();
122        assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
123    }
124
125    #[tokio::test]
126    async fn test_rejects_wrong_token() {
127        let resp = app("real-token")
128            .oneshot(
129                Request::builder()
130                    .uri("/protected")
131                    .header("authorization", "Bearer wrong-token")
132                    .body(Body::empty())
133                    .unwrap(),
134            )
135            .await
136            .unwrap();
137        assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
138    }
139
140    #[tokio::test]
141    async fn test_accepts_correct_token() {
142        let resp = app("my-token")
143            .oneshot(
144                Request::builder()
145                    .uri("/protected")
146                    .header("authorization", "Bearer my-token")
147                    .body(Body::empty())
148                    .unwrap(),
149            )
150            .await
151            .unwrap();
152        assert_eq!(resp.status(), StatusCode::OK);
153    }
154
155    #[tokio::test]
156    async fn test_accepts_lowercase_bearer_scheme() {
157        let resp = app("my-token")
158            .oneshot(
159                Request::builder()
160                    .uri("/protected")
161                    .header("authorization", "bearer my-token")
162                    .body(Body::empty())
163                    .unwrap(),
164            )
165            .await
166            .unwrap();
167        assert_eq!(resp.status(), StatusCode::OK);
168    }
169}