Skip to main content

notedthat_api_http/
middleware.rs

1//! 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::body::Body;
7use axum::extract::{MatchedPath, State};
8use axum::http::{Method, Request, StatusCode, header::WWW_AUTHENTICATE};
9use axum::middleware::Next;
10use axum::response::{IntoResponse, Response};
11use notedthat_core::{Principal, Schemes};
12use tower_http::request_id::RequestId;
13
14/// The `(method, route)` pairs an anonymous request is allowed to reach.
15///
16/// Authorization itself is per key and lives in the handlers, because a
17/// path-scoped rule cannot be evaluated from a route pattern — `read` on
18/// `{*object_path}` has no answer until the key is known. That would leave a
19/// route added later without an authorization call open to the world, so the
20/// table is inverted instead of deleted: a `(method, route)` pair absent from
21/// here is unreachable without a credential, and opening a new route means
22/// coming here and saying so.
23///
24/// Every route listed **must** have a handler that calls
25/// [`crate::authz::KbAccess::require`] or `require_any`; `route_backstop.rs`
26/// asserts the two stay in step.
27const ANONYMOUS_REACHABLE: &[(&Method, &str)] = &[
28    (&Method::GET, MATCHED_KBS),
29    (&Method::HEAD, MATCHED_KBS),
30    (&Method::GET, MATCHED_KB),
31    (&Method::HEAD, MATCHED_KB),
32    (&Method::GET, MATCHED_KB_OBJECT),
33    (&Method::HEAD, MATCHED_KB_OBJECT),
34    (&Method::POST, MATCHED_KB_SEARCH),
35];
36
37/// Axum middleware that establishes the request's [`Principal`].
38///
39/// This layer authenticates; it does not authorize. A credential the
40/// [`notedthat_core::Authenticator`] accepts makes the request
41/// [`Principal::SignedIn`], an absent credential makes it [`Principal::Anyone`],
42/// and a supplied credential that does not verify is always `401` — never
43/// quietly downgraded to anonymous, which is the rule that stops a typo'd token
44/// from silently becoming a public view.
45///
46/// Every `401` that leaves this layer — its own, or one a handler answered —
47/// carries the bearer challenge when the deployment publishes protected-
48/// resource metadata, so an MCP client can find the authorization server.
49///
50/// It is mounted on the `/api/v1` routes only. The unauthenticated root routes
51/// (`/healthz`, `/readyz`, `/llms.txt`) never reach it, and `/browse` resolves
52/// its own principal through the same authenticator because it is mounted
53/// outside this layer (see [`crate::router::browse`]).
54pub async fn auth_middleware(
55    State(state): State<AppState>,
56    mut req: Request<Body>,
57    next: Next,
58) -> Response {
59    let request_id = extract_request_id(&req);
60
61    let response = match state
62        .authenticator
63        .resolve(req.headers(), Schemes::Bearer)
64        .await
65    {
66        Err(CredentialRefused) => ApiErrorResponse::unauthorized(request_id).into_response(),
67        Ok(principal) if principal.is_anonymous() && !anonymous_may_reach(&req) => {
68            ApiErrorResponse::unauthorized(request_id).into_response()
69        }
70        Ok(principal) => {
71            req.extensions_mut().insert(principal);
72            next.run(req).await
73        }
74    };
75    with_bearer_challenge(&state, response)
76}
77
78/// Add the `WWW-Authenticate` challenge to a `401`, when there is one to add.
79pub(crate) fn with_bearer_challenge(state: &AppState, mut response: Response) -> Response {
80    if response.status() == StatusCode::UNAUTHORIZED
81        && let Some(challenge) = state.authenticator.bearer_challenge()
82    {
83        response.headers_mut().insert(WWW_AUTHENTICATE, challenge);
84    }
85    response
86}
87
88/// Whether this request's route lets an anonymous caller through to a handler
89/// that will authorize it per key.
90fn anonymous_may_reach<B>(req: &Request<B>) -> bool {
91    let Some(matched) = req.extensions().get::<MatchedPath>() else {
92        return false;
93    };
94    let matched = matched.as_str();
95    ANONYMOUS_REACHABLE
96        .iter()
97        .any(|(method, route)| *method == req.method() && *route == matched)
98}
99
100pub use notedthat_core::CredentialRefused;
101
102/// The principal established at the HTTP boundary.
103///
104/// Defaults to [`Principal::Anyone`] when the auth layer has not run, so a
105/// handler reached by an unexpected route fails closed rather than open.
106pub fn principal<B>(req: &Request<B>) -> Principal {
107    req.extensions()
108        .get::<Principal>()
109        .cloned()
110        .unwrap_or(Principal::Anyone)
111}
112
113pub use notedthat_core::is_internal_path;
114
115/// Extract the `x-request-id` value from request extensions, falling back to a
116/// generated UUID if the `SetRequestId` middleware hasn't run yet.
117pub fn extract_request_id<B>(req: &Request<B>) -> String {
118    req.extensions()
119        .get::<RequestId>()
120        .and_then(|r| r.header_value().to_str().ok())
121        .map_or_else(
122            || {
123                tracing::warn!("request_id missing from Extensions — generating fallback");
124                uuid::Uuid::now_v7().to_string()
125            },
126            str::to_string,
127        )
128}
129
130#[cfg(test)]
131mod tests {
132    use super::*;
133    use crate::testing::InMemoryStorage;
134    use axum::middleware::from_fn_with_state;
135    use axum::response::IntoResponse;
136    use axum::routing::get;
137    use axum::{Router, body::Body, http::StatusCode};
138    use std::collections::BTreeMap;
139    use std::sync::Arc;
140    use tower::util::ServiceExt;
141
142    fn test_state(token: &str) -> AppState {
143        let (indexer_tx, _rx) = tokio::sync::mpsc::channel(1024);
144        AppState {
145            storage: Arc::new(InMemoryStorage::default()),
146            declared_kbs: Arc::new(BTreeMap::new()),
147            access_policies: Arc::new(BTreeMap::new()),
148            authenticator: Arc::new(notedthat_core::Authenticator::new(token)),
149            max_body_size: 16 * 1024 * 1024,
150            max_patchable_size: 16 * 1024 * 1024,
151            indexer_tx,
152            searcher: Arc::new(crate::testing::NoopSearcher),
153        }
154    }
155
156    fn app(token: &str) -> Router {
157        let state = test_state(token);
158        Router::new()
159            .route("/protected", get(|| async { "secret".into_response() }))
160            .layer(from_fn_with_state(state.clone(), auth_middleware))
161            .with_state(state)
162    }
163
164    /// The public routes are exempt because they are mounted outside this
165    /// layer, not because the layer knows their paths. Anything that does
166    /// reach the layer without a credential is rejected — including a path
167    /// that merely looks like a probe.
168    #[tokio::test]
169    async fn no_path_is_exempt_once_the_request_reaches_this_layer() {
170        for uri in ["/healthz", "/readyz", "/llms.txt", "/protected"] {
171            let resp = app("my-token")
172                .oneshot(Request::builder().uri(uri).body(Body::empty()).unwrap())
173                .await
174                .unwrap();
175            assert_eq!(
176                resp.status(),
177                StatusCode::UNAUTHORIZED,
178                "{uri} must not be exempted by the auth layer itself"
179            );
180        }
181    }
182
183    #[tokio::test]
184    async fn test_rejects_missing_auth() {
185        let resp = app("my-token")
186            .oneshot(
187                Request::builder()
188                    .uri("/protected")
189                    .body(Body::empty())
190                    .unwrap(),
191            )
192            .await
193            .unwrap();
194        assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
195    }
196
197    #[tokio::test]
198    async fn test_rejects_wrong_token() {
199        let resp = app("real-token")
200            .oneshot(
201                Request::builder()
202                    .uri("/protected")
203                    .header("authorization", "Bearer wrong-token")
204                    .body(Body::empty())
205                    .unwrap(),
206            )
207            .await
208            .unwrap();
209        assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
210    }
211
212    #[tokio::test]
213    async fn test_accepts_correct_token() {
214        let resp = app("my-token")
215            .oneshot(
216                Request::builder()
217                    .uri("/protected")
218                    .header("authorization", "Bearer my-token")
219                    .body(Body::empty())
220                    .unwrap(),
221            )
222            .await
223            .unwrap();
224        assert_eq!(resp.status(), StatusCode::OK);
225    }
226
227    #[tokio::test]
228    async fn test_accepts_lowercase_bearer_scheme() {
229        let resp = app("my-token")
230            .oneshot(
231                Request::builder()
232                    .uri("/protected")
233                    .header("authorization", "bearer my-token")
234                    .body(Body::empty())
235                    .unwrap(),
236            )
237            .await
238            .unwrap();
239        assert_eq!(resp.status(), StatusCode::OK);
240    }
241}