notedthat_api_http/
middleware.rs1use crate::error::ApiErrorResponse;
4use crate::router::{
5 MATCHED_KB, MATCHED_KB_EVENTS, MATCHED_KB_INDEX, MATCHED_KB_INDEX_RECONCILE, MATCHED_KB_OBJECT,
6 MATCHED_KB_SEARCH, 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
93const OPERATOR_ONLY: &[&str] = &[MATCHED_KB_INDEX_RECONCILE];
97
98fn anonymous_may_reach<B>(req: &Request<B>) -> bool {
101 let Some(matched) = req.extensions().get::<MatchedPath>() else {
102 return false;
103 };
104 let matched = matched.as_str();
105 let reachable = ANONYMOUS_REACHABLE
106 .iter()
107 .any(|(method, route)| *method == req.method() && *route == matched);
108 debug_assert!(
109 !(reachable && OPERATOR_ONLY.contains(&matched)),
110 "an operator route is listed as anonymously reachable"
111 );
112 reachable
113}
114
115pub use notedthat_core::CredentialRefused;
116
117pub fn principal<B>(req: &Request<B>) -> Principal {
122 req.extensions()
123 .get::<Principal>()
124 .cloned()
125 .unwrap_or(Principal::Anyone)
126}
127
128pub use notedthat_core::is_internal_path;
129
130pub fn extract_request_id<B>(req: &Request<B>) -> String {
133 req.extensions()
134 .get::<RequestId>()
135 .and_then(|r| r.header_value().to_str().ok())
136 .map_or_else(
137 || {
138 tracing::warn!("request_id missing from Extensions — generating fallback");
139 uuid::Uuid::now_v7().to_string()
140 },
141 str::to_string,
142 )
143}
144
145#[cfg(test)]
146mod tests {
147 use super::*;
148 use crate::testing::InMemoryStorage;
149 use axum::middleware::from_fn_with_state;
150 use axum::response::IntoResponse;
151 use axum::routing::get;
152 use axum::{Router, body::Body, http::StatusCode};
153 use std::collections::BTreeMap;
154 use std::sync::Arc;
155 use tower::util::ServiceExt;
156
157 fn test_state(token: &str) -> AppState {
158 let (indexer_tx, _rx) = tokio::sync::mpsc::channel(1024);
159 AppState {
160 storage: Arc::new(InMemoryStorage::default()),
161 declared_kbs: Arc::new(BTreeMap::new()),
162 access_policies: Arc::new(BTreeMap::new()),
163 kb_details: Arc::new(BTreeMap::new()),
164 authenticator: Arc::new(notedthat_core::Authenticator::new(token)),
165 max_body_size: 16 * 1024 * 1024,
166 max_patchable_size: 16 * 1024 * 1024,
167 indexer_tx: (&indexer_tx).into(),
168 searcher: Arc::new(crate::testing::NoopSearcher),
169 events: None,
170 index_health: Arc::new(notedthat_indexer::IndexHealth::new()),
171 readiness: crate::testing::ready_receiver(),
172 reconcile: None,
173 }
174 }
175
176 fn app(token: &str) -> Router {
177 let state = test_state(token);
178 Router::new()
179 .route("/protected", get(|| async { "secret".into_response() }))
180 .layer(from_fn_with_state(state.clone(), auth_middleware))
181 .with_state(state)
182 }
183
184 #[tokio::test]
189 async fn no_path_is_exempt_once_the_request_reaches_this_layer() {
190 for uri in ["/healthz", "/readyz", "/llms.txt", "/protected"] {
191 let resp = app("my-token")
192 .oneshot(Request::builder().uri(uri).body(Body::empty()).unwrap())
193 .await
194 .unwrap();
195 assert_eq!(
196 resp.status(),
197 StatusCode::UNAUTHORIZED,
198 "{uri} must not be exempted by the auth layer itself"
199 );
200 }
201 }
202
203 #[test]
208 fn the_operator_route_is_not_anonymously_reachable() {
209 for operator in OPERATOR_ONLY {
210 assert!(
211 !ANONYMOUS_REACHABLE
212 .iter()
213 .any(|(_, route)| route == operator),
214 "{operator} must stay behind the credential check"
215 );
216 }
217 }
218
219 #[tokio::test]
220 async fn test_rejects_missing_auth() {
221 let resp = app("my-token")
222 .oneshot(
223 Request::builder()
224 .uri("/protected")
225 .body(Body::empty())
226 .unwrap(),
227 )
228 .await
229 .unwrap();
230 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
231 }
232
233 #[tokio::test]
234 async fn test_rejects_wrong_token() {
235 let resp = app("real-token")
236 .oneshot(
237 Request::builder()
238 .uri("/protected")
239 .header("authorization", "Bearer wrong-token")
240 .body(Body::empty())
241 .unwrap(),
242 )
243 .await
244 .unwrap();
245 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
246 }
247
248 #[tokio::test]
249 async fn test_accepts_correct_token() {
250 let resp = app("my-token")
251 .oneshot(
252 Request::builder()
253 .uri("/protected")
254 .header("authorization", "Bearer my-token")
255 .body(Body::empty())
256 .unwrap(),
257 )
258 .await
259 .unwrap();
260 assert_eq!(resp.status(), StatusCode::OK);
261 }
262
263 #[tokio::test]
264 async fn test_accepts_lowercase_bearer_scheme() {
265 let resp = app("my-token")
266 .oneshot(
267 Request::builder()
268 .uri("/protected")
269 .header("authorization", "bearer my-token")
270 .body(Body::empty())
271 .unwrap(),
272 )
273 .await
274 .unwrap();
275 assert_eq!(resp.status(), StatusCode::OK);
276 }
277}