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::RequestExt;
7use axum::body::Body;
8use axum::extract::{MatchedPath, Path, State};
9use axum::http::{Method, Request, header::AUTHORIZATION};
10use axum::middleware::Next;
11use axum::response::Response;
12use notedthat_core::{PublicReadCapability, extract_bearer_from_header, verify_bearer_token};
13use tower_http::request_id::RequestId;
14
15#[derive(Clone, Copy, Debug, PartialEq, Eq)]
17pub enum AuthContext {
18 Authenticated,
20 Anonymous,
22}
23
24impl AuthContext {
25 #[must_use]
27 pub const fn is_anonymous(self) -> bool {
28 matches!(self, Self::Anonymous)
29 }
30}
31
32pub 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 req.extensions_mut().insert(AuthContext::Anonymous);
63 return Ok(next.run(req).await);
64 }
65
66 Err(ApiErrorResponse::unauthorized(request_id))
67}
68
69async fn anonymous_capability(
70 req: &mut Request<Body>,
71 state: &AppState,
72) -> Option<PublicReadCapability> {
73 let matched_path = req.extensions().get::<MatchedPath>()?.as_str();
74 let capability = match (req.method(), matched_path) {
75 (&Method::GET | &Method::HEAD, MATCHED_KBS) => {
76 return state
77 .declared_kbs
78 .keys()
79 .any(|slug| {
80 state
81 .public_read_policies
82 .get(slug)
83 .is_some_and(|policy| policy.allows(PublicReadCapability::Discover))
84 })
85 .then_some(PublicReadCapability::Discover);
86 }
87 (&Method::GET | &Method::HEAD, MATCHED_KB) => PublicReadCapability::Browse,
88 (&Method::GET | &Method::HEAD, MATCHED_KB_OBJECT) => PublicReadCapability::Content,
89 (&Method::POST, MATCHED_KB_SEARCH) => PublicReadCapability::Search,
90 _ => return None,
91 };
92 let Path(params) = req
93 .extract_parts::<Path<std::collections::BTreeMap<String, String>>>()
94 .await
95 .ok()?;
96 let kb_slug = params.get("kb_slug")?;
97 state
98 .public_read_policies
99 .get(kb_slug)
100 .filter(|policy| policy.allows(capability))
101 .map(|_| capability)
102}
103
104pub fn auth_context<B>(req: &Request<B>) -> AuthContext {
106 req.extensions()
107 .get::<AuthContext>()
108 .copied()
109 .unwrap_or(AuthContext::Anonymous)
110}
111
112pub use notedthat_core::is_internal_path;
113
114pub fn extract_request_id<B>(req: &Request<B>) -> String {
117 req.extensions()
118 .get::<RequestId>()
119 .and_then(|r| r.header_value().to_str().ok())
120 .map_or_else(
121 || {
122 tracing::warn!("request_id missing from Extensions — generating fallback");
123 uuid::Uuid::now_v7().to_string()
124 },
125 str::to_string,
126 )
127}
128
129#[cfg(test)]
130mod tests {
131 use super::*;
132 use crate::testing::InMemoryStorage;
133 use axum::middleware::from_fn_with_state;
134 use axum::response::IntoResponse;
135 use axum::routing::get;
136 use axum::{Router, body::Body, http::StatusCode};
137 use std::collections::BTreeMap;
138 use std::sync::Arc;
139 use tower::util::ServiceExt;
140
141 fn test_state(token: &str) -> AppState {
142 let (indexer_tx, _rx) = tokio::sync::mpsc::channel(1024);
143 AppState {
144 storage: Arc::new(InMemoryStorage::default()),
145 declared_kbs: Arc::new(BTreeMap::new()),
146 public_read_policies: Arc::new(BTreeMap::new()),
147 bearer_token: Arc::new(token.to_string()),
148 max_body_size: 16 * 1024 * 1024,
149 max_patchable_size: 16 * 1024 * 1024,
150 indexer_tx,
151 searcher: Arc::new(crate::testing::NoopSearcher),
152 }
153 }
154
155 fn app(token: &str) -> Router {
156 let state = test_state(token);
157 Router::new()
158 .route("/protected", get(|| async { "secret".into_response() }))
159 .layer(from_fn_with_state(state.clone(), auth_middleware))
160 .with_state(state)
161 }
162
163 #[tokio::test]
168 async fn no_path_is_exempt_once_the_request_reaches_this_layer() {
169 for uri in ["/healthz", "/readyz", "/llms.txt", "/protected"] {
170 let resp = app("my-token")
171 .oneshot(Request::builder().uri(uri).body(Body::empty()).unwrap())
172 .await
173 .unwrap();
174 assert_eq!(
175 resp.status(),
176 StatusCode::UNAUTHORIZED,
177 "{uri} must not be exempted by the auth layer itself"
178 );
179 }
180 }
181
182 #[tokio::test]
183 async fn test_rejects_missing_auth() {
184 let resp = app("my-token")
185 .oneshot(
186 Request::builder()
187 .uri("/protected")
188 .body(Body::empty())
189 .unwrap(),
190 )
191 .await
192 .unwrap();
193 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
194 }
195
196 #[tokio::test]
197 async fn test_rejects_wrong_token() {
198 let resp = app("real-token")
199 .oneshot(
200 Request::builder()
201 .uri("/protected")
202 .header("authorization", "Bearer wrong-token")
203 .body(Body::empty())
204 .unwrap(),
205 )
206 .await
207 .unwrap();
208 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
209 }
210
211 #[tokio::test]
212 async fn test_accepts_correct_token() {
213 let resp = app("my-token")
214 .oneshot(
215 Request::builder()
216 .uri("/protected")
217 .header("authorization", "Bearer my-token")
218 .body(Body::empty())
219 .unwrap(),
220 )
221 .await
222 .unwrap();
223 assert_eq!(resp.status(), StatusCode::OK);
224 }
225
226 #[tokio::test]
227 async fn test_accepts_lowercase_bearer_scheme() {
228 let resp = app("my-token")
229 .oneshot(
230 Request::builder()
231 .uri("/protected")
232 .header("authorization", "bearer my-token")
233 .body(Body::empty())
234 .unwrap(),
235 )
236 .await
237 .unwrap();
238 assert_eq!(resp.status(), StatusCode::OK);
239 }
240}