notedthat_api_http/
middleware.rs1use crate::error::ApiErrorResponse;
4use crate::router::{
5 MATCHED_KB, MATCHED_KB_EVENTS, MATCHED_KB_INDEX, MATCHED_KB_OBJECT, MATCHED_KB_SEARCH,
6 MATCHED_KBS,
7};
8use crate::state::AppState;
9use axum::body::Body;
10use axum::extract::{MatchedPath, State};
11use axum::http::{Method, Request, StatusCode, header::WWW_AUTHENTICATE};
12use axum::middleware::Next;
13use axum::response::{IntoResponse, Response};
14use notedthat_core::{Principal, Schemes};
15use tower_http::request_id::RequestId;
16
17const ANONYMOUS_REACHABLE: &[(&Method, &str)] = &[
31 (&Method::GET, MATCHED_KBS),
32 (&Method::HEAD, MATCHED_KBS),
33 (&Method::GET, MATCHED_KB),
34 (&Method::HEAD, MATCHED_KB),
35 (&Method::GET, MATCHED_KB_OBJECT),
36 (&Method::HEAD, MATCHED_KB_OBJECT),
37 (&Method::POST, MATCHED_KB_SEARCH),
38 (&Method::GET, MATCHED_KB_EVENTS),
39 (&Method::GET, MATCHED_KB_INDEX),
40];
41
42pub async fn auth_middleware(
60 State(state): State<AppState>,
61 mut req: Request<Body>,
62 next: Next,
63) -> Response {
64 let request_id = extract_request_id(&req);
65
66 let response = match state
67 .authenticator
68 .resolve(req.headers(), Schemes::Bearer)
69 .await
70 {
71 Err(CredentialRefused) => ApiErrorResponse::unauthorized(request_id).into_response(),
72 Ok(principal) if principal.is_anonymous() && !anonymous_may_reach(&req) => {
73 ApiErrorResponse::unauthorized(request_id).into_response()
74 }
75 Ok(principal) => {
76 req.extensions_mut().insert(principal);
77 next.run(req).await
78 }
79 };
80 with_bearer_challenge(&state, response)
81}
82
83pub(crate) fn with_bearer_challenge(state: &AppState, mut response: Response) -> Response {
85 if response.status() == StatusCode::UNAUTHORIZED
86 && let Some(challenge) = state.authenticator.bearer_challenge()
87 {
88 response.headers_mut().insert(WWW_AUTHENTICATE, challenge);
89 }
90 response
91}
92
93fn anonymous_may_reach<B>(req: &Request<B>) -> bool {
96 let Some(matched) = req.extensions().get::<MatchedPath>() else {
97 return false;
98 };
99 let matched = matched.as_str();
100 ANONYMOUS_REACHABLE
101 .iter()
102 .any(|(method, route)| *method == req.method() && *route == matched)
103}
104
105pub use notedthat_core::CredentialRefused;
106
107pub fn principal<B>(req: &Request<B>) -> Principal {
112 req.extensions()
113 .get::<Principal>()
114 .cloned()
115 .unwrap_or(Principal::Anyone)
116}
117
118pub use notedthat_core::is_internal_path;
119
120pub 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 access_policies: Arc::new(BTreeMap::new()),
153 kb_details: Arc::new(BTreeMap::new()),
154 authenticator: Arc::new(notedthat_core::Authenticator::new(token)),
155 max_body_size: 16 * 1024 * 1024,
156 max_patchable_size: 16 * 1024 * 1024,
157 indexer_tx,
158 searcher: Arc::new(crate::testing::NoopSearcher),
159 events: None,
160 index_health: Arc::new(notedthat_indexer::IndexHealth::new()),
161 }
162 }
163
164 fn app(token: &str) -> Router {
165 let state = test_state(token);
166 Router::new()
167 .route("/protected", get(|| async { "secret".into_response() }))
168 .layer(from_fn_with_state(state.clone(), auth_middleware))
169 .with_state(state)
170 }
171
172 #[tokio::test]
177 async fn no_path_is_exempt_once_the_request_reaches_this_layer() {
178 for uri in ["/healthz", "/readyz", "/llms.txt", "/protected"] {
179 let resp = app("my-token")
180 .oneshot(Request::builder().uri(uri).body(Body::empty()).unwrap())
181 .await
182 .unwrap();
183 assert_eq!(
184 resp.status(),
185 StatusCode::UNAUTHORIZED,
186 "{uri} must not be exempted by the auth layer itself"
187 );
188 }
189 }
190
191 #[tokio::test]
192 async fn test_rejects_missing_auth() {
193 let resp = app("my-token")
194 .oneshot(
195 Request::builder()
196 .uri("/protected")
197 .body(Body::empty())
198 .unwrap(),
199 )
200 .await
201 .unwrap();
202 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
203 }
204
205 #[tokio::test]
206 async fn test_rejects_wrong_token() {
207 let resp = app("real-token")
208 .oneshot(
209 Request::builder()
210 .uri("/protected")
211 .header("authorization", "Bearer wrong-token")
212 .body(Body::empty())
213 .unwrap(),
214 )
215 .await
216 .unwrap();
217 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
218 }
219
220 #[tokio::test]
221 async fn test_accepts_correct_token() {
222 let resp = app("my-token")
223 .oneshot(
224 Request::builder()
225 .uri("/protected")
226 .header("authorization", "Bearer my-token")
227 .body(Body::empty())
228 .unwrap(),
229 )
230 .await
231 .unwrap();
232 assert_eq!(resp.status(), StatusCode::OK);
233 }
234
235 #[tokio::test]
236 async fn test_accepts_lowercase_bearer_scheme() {
237 let resp = app("my-token")
238 .oneshot(
239 Request::builder()
240 .uri("/protected")
241 .header("authorization", "bearer my-token")
242 .body(Body::empty())
243 .unwrap(),
244 )
245 .await
246 .unwrap();
247 assert_eq!(resp.status(), StatusCode::OK);
248 }
249}