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::RequestExt;
6use axum::body::Body;
7use axum::extract::{MatchedPath, Path, State};
8use axum::http::{Method, Request, header::AUTHORIZATION};
9use axum::middleware::Next;
10use axum::response::Response;
11use notedthat_core::{PublicReadCapability, extract_bearer_from_header, verify_bearer_token};
12use tower_http::request_id::RequestId;
13
14/// Root-level paths that bypass Bearer authentication.
15const AUTH_EXEMPT_PATHS: &[&str] = &["/healthz", "/readyz", "/llms.txt"];
16
17/// Authentication state established at the HTTP boundary.
18#[derive(Clone, Copy, Debug, PartialEq, Eq)]
19pub enum AuthContext {
20    /// A valid Bearer token was supplied.
21    Authenticated,
22    /// No credentials were supplied and a public route capability allowed access.
23    Anonymous,
24}
25
26impl AuthContext {
27    /// Return whether the request is using anonymous public-read access.
28    #[must_use]
29    pub const fn is_anonymous(self) -> bool {
30        matches!(self, Self::Anonymous)
31    }
32}
33
34/// Axum middleware that validates the `Authorization: Bearer <token>` header.
35///
36/// A supplied credential must always be a single valid Bearer token. When no
37/// credential is supplied, only root public paths and explicitly granted read
38/// capabilities pass through as anonymous requests.
39pub async fn auth_middleware(
40    State(state): State<AppState>,
41    mut req: Request<Body>,
42    next: Next,
43) -> Result<Response, ApiErrorResponse> {
44    let request_id = extract_request_id(&req);
45
46    let mut authorization_values = req.headers().get_all(AUTHORIZATION).iter();
47    if let Some(header) = authorization_values.next() {
48        if authorization_values.next().is_some() {
49            return Err(ApiErrorResponse::unauthorized(request_id));
50        }
51        header
52            .to_str()
53            .ok()
54            .and_then(extract_bearer_from_header)
55            .filter(|token| verify_bearer_token(token, &state.bearer_token))
56            .ok_or_else(|| ApiErrorResponse::unauthorized(request_id.clone()))?;
57        req.extensions_mut().insert(AuthContext::Authenticated);
58        return Ok(next.run(req).await);
59    }
60
61    if anonymous_capability(&mut req, &state).await.is_some()
62        || AUTH_EXEMPT_PATHS.contains(&req.uri().path())
63    {
64        req.extensions_mut().insert(AuthContext::Anonymous);
65        return Ok(next.run(req).await);
66    }
67
68    Err(ApiErrorResponse::unauthorized(request_id))
69}
70
71async fn anonymous_capability(
72    req: &mut Request<Body>,
73    state: &AppState,
74) -> Option<PublicReadCapability> {
75    let matched_path = req.extensions().get::<MatchedPath>()?.as_str();
76    let capability = match (req.method(), matched_path) {
77        (&Method::GET | &Method::HEAD, "/v1/knowledgebases") => {
78            return state
79                .declared_kbs
80                .keys()
81                .any(|slug| {
82                    state
83                        .public_read_policies
84                        .get(slug)
85                        .is_some_and(|policy| policy.allows(PublicReadCapability::Discover))
86                })
87                .then_some(PublicReadCapability::Discover);
88        }
89        (&Method::GET | &Method::HEAD, "/v1/knowledgebases/{kb_slug}") => {
90            PublicReadCapability::Browse
91        }
92        (&Method::GET | &Method::HEAD, "/v1/knowledgebases/{kb_slug}/{*object_path}") => {
93            PublicReadCapability::Content
94        }
95        (&Method::POST, "/v1/knowledgebases/{kb_slug}/search") => PublicReadCapability::Search,
96        _ => return None,
97    };
98    let Path(params) = req
99        .extract_parts::<Path<std::collections::BTreeMap<String, String>>>()
100        .await
101        .ok()?;
102    let kb_slug = params.get("kb_slug")?;
103    state
104        .public_read_policies
105        .get(kb_slug)
106        .filter(|policy| policy.allows(capability))
107        .map(|_| capability)
108}
109
110/// Return the request authentication context, defaulting to anonymous when the auth layer has not run.
111pub fn auth_context<B>(req: &Request<B>) -> AuthContext {
112    req.extensions()
113        .get::<AuthContext>()
114        .copied()
115        .unwrap_or(AuthContext::Anonymous)
116}
117
118pub use notedthat_core::is_internal_path;
119
120/// Extract the `x-request-id` value from request extensions, falling back to a
121/// generated UUID if the `SetRequestId` middleware hasn't run yet.
122pub fn extract_request_id<B>(req: &Request<B>) -> String {
123    req.extensions()
124        .get::<RequestId>()
125        .and_then(|r| r.header_value().to_str().ok())
126        .map_or_else(
127            || {
128                tracing::warn!("request_id missing from Extensions — generating fallback");
129                uuid::Uuid::now_v7().to_string()
130            },
131            str::to_string,
132        )
133}
134
135#[cfg(test)]
136mod tests {
137    use super::*;
138    use crate::testing::InMemoryStorage;
139    use axum::middleware::from_fn_with_state;
140    use axum::response::IntoResponse;
141    use axum::routing::get;
142    use axum::{Router, body::Body, http::StatusCode};
143    use std::collections::BTreeMap;
144    use std::sync::Arc;
145    use tower::util::ServiceExt;
146
147    fn test_state(token: &str) -> AppState {
148        let (indexer_tx, _rx) = tokio::sync::mpsc::channel(1024);
149        AppState {
150            storage: Arc::new(InMemoryStorage::default()),
151            declared_kbs: Arc::new(BTreeMap::new()),
152            public_read_policies: Arc::new(BTreeMap::new()),
153            bearer_token: Arc::new(token.to_string()),
154            max_body_size: 16 * 1024 * 1024,
155            max_patchable_size: 16 * 1024 * 1024,
156            indexer_tx,
157            searcher: Arc::new(crate::testing::NoopSearcher),
158        }
159    }
160
161    fn app(token: &str) -> Router {
162        let state = test_state(token);
163        Router::new()
164            .route("/healthz", get(|| async { "ok" }))
165            .route("/protected", get(|| async { "secret".into_response() }))
166            .layer(from_fn_with_state(state.clone(), auth_middleware))
167            .with_state(state)
168    }
169
170    #[tokio::test]
171    async fn test_healthz_bypasses_auth() {
172        let resp = app("my-token")
173            .oneshot(
174                Request::builder()
175                    .uri("/healthz")
176                    .body(Body::empty())
177                    .unwrap(),
178            )
179            .await
180            .unwrap();
181        assert_eq!(resp.status(), StatusCode::OK);
182    }
183
184    #[tokio::test]
185    async fn test_rejects_missing_auth() {
186        let resp = app("my-token")
187            .oneshot(
188                Request::builder()
189                    .uri("/protected")
190                    .body(Body::empty())
191                    .unwrap(),
192            )
193            .await
194            .unwrap();
195        assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
196    }
197
198    #[tokio::test]
199    async fn test_rejects_wrong_token() {
200        let resp = app("real-token")
201            .oneshot(
202                Request::builder()
203                    .uri("/protected")
204                    .header("authorization", "Bearer wrong-token")
205                    .body(Body::empty())
206                    .unwrap(),
207            )
208            .await
209            .unwrap();
210        assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
211    }
212
213    #[tokio::test]
214    async fn test_accepts_correct_token() {
215        let resp = app("my-token")
216            .oneshot(
217                Request::builder()
218                    .uri("/protected")
219                    .header("authorization", "Bearer my-token")
220                    .body(Body::empty())
221                    .unwrap(),
222            )
223            .await
224            .unwrap();
225        assert_eq!(resp.status(), StatusCode::OK);
226    }
227
228    #[tokio::test]
229    async fn test_accepts_lowercase_bearer_scheme() {
230        let resp = app("my-token")
231            .oneshot(
232                Request::builder()
233                    .uri("/protected")
234                    .header("authorization", "bearer my-token")
235                    .body(Body::empty())
236                    .unwrap(),
237            )
238            .await
239            .unwrap();
240        assert_eq!(resp.status(), StatusCode::OK);
241    }
242}