notedthat_api_http/
middleware.rs1use 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
14const 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
37pub 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
78pub(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
88fn 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
102pub 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
115pub 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 #[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}