notedthat_api_http/
middleware.rs1use crate::error::ApiErrorResponse;
4use crate::state::AppState;
5use axum::RequestExt;
6use axum::body::Body;
7use axum::extract::{MatchedPath, Path, State};
8use axum::http::{Method, Request, header::AUTHORIZATION};
9use axum::middleware::Next;
10use axum::response::Response;
11use notedthat_core::{PublicReadCapability, extract_bearer_from_header, verify_bearer_token};
12use tower_http::request_id::RequestId;
13
14const AUTH_EXEMPT_PATHS: &[&str] = &["/healthz", "/readyz", "/llms.txt"];
16
17#[derive(Clone, Copy, Debug, PartialEq, Eq)]
19pub enum AuthContext {
20 Authenticated,
22 Anonymous,
24}
25
26impl AuthContext {
27 #[must_use]
29 pub const fn is_anonymous(self) -> bool {
30 matches!(self, Self::Anonymous)
31 }
32}
33
34pub async fn auth_middleware(
40 State(state): State<AppState>,
41 mut req: Request<Body>,
42 next: Next,
43) -> Result<Response, ApiErrorResponse> {
44 let request_id = extract_request_id(&req);
45
46 let mut authorization_values = req.headers().get_all(AUTHORIZATION).iter();
47 if let Some(header) = authorization_values.next() {
48 if authorization_values.next().is_some() {
49 return Err(ApiErrorResponse::unauthorized(request_id));
50 }
51 header
52 .to_str()
53 .ok()
54 .and_then(extract_bearer_from_header)
55 .filter(|token| verify_bearer_token(token, &state.bearer_token))
56 .ok_or_else(|| ApiErrorResponse::unauthorized(request_id.clone()))?;
57 req.extensions_mut().insert(AuthContext::Authenticated);
58 return Ok(next.run(req).await);
59 }
60
61 if anonymous_capability(&mut req, &state).await.is_some()
62 || AUTH_EXEMPT_PATHS.contains(&req.uri().path())
63 {
64 req.extensions_mut().insert(AuthContext::Anonymous);
65 return Ok(next.run(req).await);
66 }
67
68 Err(ApiErrorResponse::unauthorized(request_id))
69}
70
71async fn anonymous_capability(
72 req: &mut Request<Body>,
73 state: &AppState,
74) -> Option<PublicReadCapability> {
75 let matched_path = req.extensions().get::<MatchedPath>()?.as_str();
76 let capability = match (req.method(), matched_path) {
77 (&Method::GET | &Method::HEAD, "/v1/knowledgebases") => {
78 return state
79 .declared_kbs
80 .keys()
81 .any(|slug| {
82 state
83 .public_read_policies
84 .get(slug)
85 .is_some_and(|policy| policy.allows(PublicReadCapability::Discover))
86 })
87 .then_some(PublicReadCapability::Discover);
88 }
89 (&Method::GET | &Method::HEAD, "/v1/knowledgebases/{kb_slug}") => {
90 PublicReadCapability::Browse
91 }
92 (&Method::GET | &Method::HEAD, "/v1/knowledgebases/{kb_slug}/{*object_path}") => {
93 PublicReadCapability::Content
94 }
95 (&Method::POST, "/v1/knowledgebases/{kb_slug}/search") => PublicReadCapability::Search,
96 _ => return None,
97 };
98 let Path(params) = req
99 .extract_parts::<Path<std::collections::BTreeMap<String, String>>>()
100 .await
101 .ok()?;
102 let kb_slug = params.get("kb_slug")?;
103 state
104 .public_read_policies
105 .get(kb_slug)
106 .filter(|policy| policy.allows(capability))
107 .map(|_| capability)
108}
109
110pub fn auth_context<B>(req: &Request<B>) -> AuthContext {
112 req.extensions()
113 .get::<AuthContext>()
114 .copied()
115 .unwrap_or(AuthContext::Anonymous)
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 public_read_policies: Arc::new(BTreeMap::new()),
153 bearer_token: Arc::new(token.to_string()),
154 max_body_size: 16 * 1024 * 1024,
155 max_patchable_size: 16 * 1024 * 1024,
156 indexer_tx,
157 searcher: Arc::new(crate::testing::NoopSearcher),
158 }
159 }
160
161 fn app(token: &str) -> Router {
162 let state = test_state(token);
163 Router::new()
164 .route("/healthz", get(|| async { "ok" }))
165 .route("/protected", get(|| async { "secret".into_response() }))
166 .layer(from_fn_with_state(state.clone(), auth_middleware))
167 .with_state(state)
168 }
169
170 #[tokio::test]
171 async fn test_healthz_bypasses_auth() {
172 let resp = app("my-token")
173 .oneshot(
174 Request::builder()
175 .uri("/healthz")
176 .body(Body::empty())
177 .unwrap(),
178 )
179 .await
180 .unwrap();
181 assert_eq!(resp.status(), StatusCode::OK);
182 }
183
184 #[tokio::test]
185 async fn test_rejects_missing_auth() {
186 let resp = app("my-token")
187 .oneshot(
188 Request::builder()
189 .uri("/protected")
190 .body(Body::empty())
191 .unwrap(),
192 )
193 .await
194 .unwrap();
195 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
196 }
197
198 #[tokio::test]
199 async fn test_rejects_wrong_token() {
200 let resp = app("real-token")
201 .oneshot(
202 Request::builder()
203 .uri("/protected")
204 .header("authorization", "Bearer wrong-token")
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_accepts_correct_token() {
215 let resp = app("my-token")
216 .oneshot(
217 Request::builder()
218 .uri("/protected")
219 .header("authorization", "Bearer my-token")
220 .body(Body::empty())
221 .unwrap(),
222 )
223 .await
224 .unwrap();
225 assert_eq!(resp.status(), StatusCode::OK);
226 }
227
228 #[tokio::test]
229 async fn test_accepts_lowercase_bearer_scheme() {
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}