notedthat_api_http/
middleware.rs1use crate::error::ApiErrorResponse;
4use crate::router::{
5 MATCHED_KB, MATCHED_KB_EVENTS, MATCHED_KB_OBJECT, MATCHED_KB_SEARCH, MATCHED_KBS,
6};
7use crate::state::AppState;
8use axum::body::Body;
9use axum::extract::{MatchedPath, State};
10use axum::http::{Method, Request, StatusCode, header::WWW_AUTHENTICATE};
11use axum::middleware::Next;
12use axum::response::{IntoResponse, Response};
13use notedthat_core::{Principal, Schemes};
14use tower_http::request_id::RequestId;
15
16const ANONYMOUS_REACHABLE: &[(&Method, &str)] = &[
30 (&Method::GET, MATCHED_KBS),
31 (&Method::HEAD, MATCHED_KBS),
32 (&Method::GET, MATCHED_KB),
33 (&Method::HEAD, MATCHED_KB),
34 (&Method::GET, MATCHED_KB_OBJECT),
35 (&Method::HEAD, MATCHED_KB_OBJECT),
36 (&Method::POST, MATCHED_KB_SEARCH),
37 (&Method::GET, MATCHED_KB_EVENTS),
38];
39
40pub async fn auth_middleware(
58 State(state): State<AppState>,
59 mut req: Request<Body>,
60 next: Next,
61) -> Response {
62 let request_id = extract_request_id(&req);
63
64 let response = match state
65 .authenticator
66 .resolve(req.headers(), Schemes::Bearer)
67 .await
68 {
69 Err(CredentialRefused) => ApiErrorResponse::unauthorized(request_id).into_response(),
70 Ok(principal) if principal.is_anonymous() && !anonymous_may_reach(&req) => {
71 ApiErrorResponse::unauthorized(request_id).into_response()
72 }
73 Ok(principal) => {
74 req.extensions_mut().insert(principal);
75 next.run(req).await
76 }
77 };
78 with_bearer_challenge(&state, response)
79}
80
81pub(crate) fn with_bearer_challenge(state: &AppState, mut response: Response) -> Response {
83 if response.status() == StatusCode::UNAUTHORIZED
84 && let Some(challenge) = state.authenticator.bearer_challenge()
85 {
86 response.headers_mut().insert(WWW_AUTHENTICATE, challenge);
87 }
88 response
89}
90
91fn anonymous_may_reach<B>(req: &Request<B>) -> bool {
94 let Some(matched) = req.extensions().get::<MatchedPath>() else {
95 return false;
96 };
97 let matched = matched.as_str();
98 ANONYMOUS_REACHABLE
99 .iter()
100 .any(|(method, route)| *method == req.method() && *route == matched)
101}
102
103pub use notedthat_core::CredentialRefused;
104
105pub fn principal<B>(req: &Request<B>) -> Principal {
110 req.extensions()
111 .get::<Principal>()
112 .cloned()
113 .unwrap_or(Principal::Anyone)
114}
115
116pub use notedthat_core::is_internal_path;
117
118pub fn extract_request_id<B>(req: &Request<B>) -> String {
121 req.extensions()
122 .get::<RequestId>()
123 .and_then(|r| r.header_value().to_str().ok())
124 .map_or_else(
125 || {
126 tracing::warn!("request_id missing from Extensions — generating fallback");
127 uuid::Uuid::now_v7().to_string()
128 },
129 str::to_string,
130 )
131}
132
133#[cfg(test)]
134mod tests {
135 use super::*;
136 use crate::testing::InMemoryStorage;
137 use axum::middleware::from_fn_with_state;
138 use axum::response::IntoResponse;
139 use axum::routing::get;
140 use axum::{Router, body::Body, http::StatusCode};
141 use std::collections::BTreeMap;
142 use std::sync::Arc;
143 use tower::util::ServiceExt;
144
145 fn test_state(token: &str) -> AppState {
146 let (indexer_tx, _rx) = tokio::sync::mpsc::channel(1024);
147 AppState {
148 storage: Arc::new(InMemoryStorage::default()),
149 declared_kbs: Arc::new(BTreeMap::new()),
150 access_policies: Arc::new(BTreeMap::new()),
151 authenticator: Arc::new(notedthat_core::Authenticator::new(token)),
152 max_body_size: 16 * 1024 * 1024,
153 max_patchable_size: 16 * 1024 * 1024,
154 indexer_tx,
155 searcher: Arc::new(crate::testing::NoopSearcher),
156 events: None,
157 }
158 }
159
160 fn app(token: &str) -> Router {
161 let state = test_state(token);
162 Router::new()
163 .route("/protected", get(|| async { "secret".into_response() }))
164 .layer(from_fn_with_state(state.clone(), auth_middleware))
165 .with_state(state)
166 }
167
168 #[tokio::test]
173 async fn no_path_is_exempt_once_the_request_reaches_this_layer() {
174 for uri in ["/healthz", "/readyz", "/llms.txt", "/protected"] {
175 let resp = app("my-token")
176 .oneshot(Request::builder().uri(uri).body(Body::empty()).unwrap())
177 .await
178 .unwrap();
179 assert_eq!(
180 resp.status(),
181 StatusCode::UNAUTHORIZED,
182 "{uri} must not be exempted by the auth layer itself"
183 );
184 }
185 }
186
187 #[tokio::test]
188 async fn test_rejects_missing_auth() {
189 let resp = app("my-token")
190 .oneshot(
191 Request::builder()
192 .uri("/protected")
193 .body(Body::empty())
194 .unwrap(),
195 )
196 .await
197 .unwrap();
198 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
199 }
200
201 #[tokio::test]
202 async fn test_rejects_wrong_token() {
203 let resp = app("real-token")
204 .oneshot(
205 Request::builder()
206 .uri("/protected")
207 .header("authorization", "Bearer wrong-token")
208 .body(Body::empty())
209 .unwrap(),
210 )
211 .await
212 .unwrap();
213 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
214 }
215
216 #[tokio::test]
217 async fn test_accepts_correct_token() {
218 let resp = app("my-token")
219 .oneshot(
220 Request::builder()
221 .uri("/protected")
222 .header("authorization", "Bearer my-token")
223 .body(Body::empty())
224 .unwrap(),
225 )
226 .await
227 .unwrap();
228 assert_eq!(resp.status(), StatusCode::OK);
229 }
230
231 #[tokio::test]
232 async fn test_accepts_lowercase_bearer_scheme() {
233 let resp = app("my-token")
234 .oneshot(
235 Request::builder()
236 .uri("/protected")
237 .header("authorization", "bearer my-token")
238 .body(Body::empty())
239 .unwrap(),
240 )
241 .await
242 .unwrap();
243 assert_eq!(resp.status(), StatusCode::OK);
244 }
245}