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, header::AUTHORIZATION};
9use axum::middleware::Next;
10use axum::response::Response;
11use notedthat_core::{Principal, extract_bearer_from_header, verify_bearer_token};
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(
50 State(state): State<AppState>,
51 mut req: Request<Body>,
52 next: Next,
53) -> Result<Response, ApiErrorResponse> {
54 let request_id = extract_request_id(&req);
55
56 let principal = resolve_principal(req.headers(), &state.bearer_token)
57 .map_err(|CredentialRefused| ApiErrorResponse::unauthorized(request_id.clone()))?;
58
59 if principal == Principal::Anyone && !anonymous_may_reach(&req) {
60 return Err(ApiErrorResponse::unauthorized(request_id));
61 }
62 req.extensions_mut().insert(principal);
63 Ok(next.run(req).await)
64}
65
66fn anonymous_may_reach<B>(req: &Request<B>) -> bool {
69 let Some(matched) = req.extensions().get::<MatchedPath>() else {
70 return false;
71 };
72 let matched = matched.as_str();
73 ANONYMOUS_REACHABLE
74 .iter()
75 .any(|(method, route)| *method == req.method() && *route == matched)
76}
77
78pub fn resolve_principal(
91 headers: &axum::http::HeaderMap,
92 expected_token: &str,
93) -> Result<Principal, CredentialRefused> {
94 let mut values = headers.get_all(AUTHORIZATION).iter();
95 let Some(header) = values.next() else {
96 return Ok(Principal::Anyone);
97 };
98 if values.next().is_some() {
99 return Err(CredentialRefused);
100 }
101 header
102 .to_str()
103 .ok()
104 .and_then(extract_bearer_from_header)
105 .filter(|token| verify_bearer_token(token, expected_token))
106 .map(|_| Principal::SignedIn)
107 .ok_or(CredentialRefused)
108}
109
110#[derive(Debug, Clone, Copy, PartialEq, Eq)]
116pub struct CredentialRefused;
117
118pub fn principal<B>(req: &Request<B>) -> Principal {
123 req.extensions()
124 .get::<Principal>()
125 .copied()
126 .unwrap_or(Principal::Anyone)
127}
128
129pub use notedthat_core::is_internal_path;
130
131pub fn extract_request_id<B>(req: &Request<B>) -> String {
134 req.extensions()
135 .get::<RequestId>()
136 .and_then(|r| r.header_value().to_str().ok())
137 .map_or_else(
138 || {
139 tracing::warn!("request_id missing from Extensions — generating fallback");
140 uuid::Uuid::now_v7().to_string()
141 },
142 str::to_string,
143 )
144}
145
146#[cfg(test)]
147mod tests {
148 use super::*;
149 use crate::testing::InMemoryStorage;
150 use axum::middleware::from_fn_with_state;
151 use axum::response::IntoResponse;
152 use axum::routing::get;
153 use axum::{Router, body::Body, http::StatusCode};
154 use std::collections::BTreeMap;
155 use std::sync::Arc;
156 use tower::util::ServiceExt;
157
158 fn test_state(token: &str) -> AppState {
159 let (indexer_tx, _rx) = tokio::sync::mpsc::channel(1024);
160 AppState {
161 storage: Arc::new(InMemoryStorage::default()),
162 declared_kbs: Arc::new(BTreeMap::new()),
163 access_policies: Arc::new(BTreeMap::new()),
164 bearer_token: Arc::new(token.to_string()),
165 max_body_size: 16 * 1024 * 1024,
166 max_patchable_size: 16 * 1024 * 1024,
167 indexer_tx,
168 searcher: Arc::new(crate::testing::NoopSearcher),
169 }
170 }
171
172 fn app(token: &str) -> Router {
173 let state = test_state(token);
174 Router::new()
175 .route("/protected", get(|| async { "secret".into_response() }))
176 .layer(from_fn_with_state(state.clone(), auth_middleware))
177 .with_state(state)
178 }
179
180 #[tokio::test]
185 async fn no_path_is_exempt_once_the_request_reaches_this_layer() {
186 for uri in ["/healthz", "/readyz", "/llms.txt", "/protected"] {
187 let resp = app("my-token")
188 .oneshot(Request::builder().uri(uri).body(Body::empty()).unwrap())
189 .await
190 .unwrap();
191 assert_eq!(
192 resp.status(),
193 StatusCode::UNAUTHORIZED,
194 "{uri} must not be exempted by the auth layer itself"
195 );
196 }
197 }
198
199 #[tokio::test]
200 async fn test_rejects_missing_auth() {
201 let resp = app("my-token")
202 .oneshot(
203 Request::builder()
204 .uri("/protected")
205 .body(Body::empty())
206 .unwrap(),
207 )
208 .await
209 .unwrap();
210 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
211 }
212
213 #[tokio::test]
214 async fn test_rejects_wrong_token() {
215 let resp = app("real-token")
216 .oneshot(
217 Request::builder()
218 .uri("/protected")
219 .header("authorization", "Bearer wrong-token")
220 .body(Body::empty())
221 .unwrap(),
222 )
223 .await
224 .unwrap();
225 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
226 }
227
228 #[tokio::test]
229 async fn test_accepts_correct_token() {
230 let resp = app("my-token")
231 .oneshot(
232 Request::builder()
233 .uri("/protected")
234 .header("authorization", "Bearer my-token")
235 .body(Body::empty())
236 .unwrap(),
237 )
238 .await
239 .unwrap();
240 assert_eq!(resp.status(), StatusCode::OK);
241 }
242
243 #[tokio::test]
244 async fn test_accepts_lowercase_bearer_scheme() {
245 let resp = app("my-token")
246 .oneshot(
247 Request::builder()
248 .uri("/protected")
249 .header("authorization", "bearer my-token")
250 .body(Body::empty())
251 .unwrap(),
252 )
253 .await
254 .unwrap();
255 assert_eq!(resp.status(), StatusCode::OK);
256 }
257}