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