1use std::sync::Arc;
7
8use axum::extract::Request;
9use axum::http::{StatusCode, header};
10use axum::middleware::Next;
11use axum::response::{IntoResponse, Response};
12use subtle::ConstantTimeEq;
13
14use koan_core::auth;
15use koan_core::db::pool::Pool;
16
17use super::AuthUser;
18
19#[derive(Clone)]
21pub struct AuthState {
22 pub public_pem: Arc<Vec<u8>>,
24 pub auth_enabled: bool,
26 pub introspection_key: Option<Arc<String>>,
29 pub pool: Arc<Pool>,
31}
32
33pub async fn auth_middleware(
38 axum::extract::State(state): axum::extract::State<AuthState>,
39 mut request: Request,
40 next: Next,
41) -> Response {
42 if !state.auth_enabled {
43 request.extensions_mut().insert(AuthUser::anonymous_admin());
44 return next.run(request).await;
45 }
46
47 if let Some(ref expected_key) = state.introspection_key
49 && let Some(provided) = request
50 .headers()
51 .get("X-Introspection-Key")
52 .and_then(|v| v.to_str().ok())
53 && provided
54 .as_bytes()
55 .ct_eq(expected_key.as_bytes())
56 .unwrap_u8()
57 == 1
58 {
59 request.extensions_mut().insert(AuthUser::anonymous_admin());
60 return next.run(request).await;
61 }
62
63 let Some(token) = extract_token(&request) else {
64 return (
65 StatusCode::UNAUTHORIZED,
66 [("WWW-Authenticate", "Bearer")],
67 "missing or invalid Authorization header",
68 )
69 .into_response();
70 };
71
72 let mark = auth::account_mark();
73 let user = match auth::validate_access_token(&state.public_pem, &token) {
74 Ok(claims) => {
75 let expires = claims.exp;
76 super::current_user(&state.pool, claims)
77 .await
78 .map(|user| (user, expires))
79 }
80 Err(_) => None,
81 };
82 match user {
83 Some((user, expires)) => {
84 request.extensions_mut().insert(super::Lease {
85 user_id: user.user_id,
86 mark,
87 expires: Some(expires),
88 });
89 request.extensions_mut().insert(user);
90 next.run(request).await
91 }
92 None => (
93 StatusCode::UNAUTHORIZED,
94 [("WWW-Authenticate", "Bearer")],
95 "invalid or expired token",
96 )
97 .into_response(),
98 }
99}
100
101fn extract_token(request: &Request) -> Option<String> {
107 request
108 .headers()
109 .get(header::COOKIE)
110 .and_then(|v| v.to_str().ok())
111 .and_then(|cookies| {
112 cookies
113 .split(';')
114 .find_map(|c| c.trim().strip_prefix("koan_access=").map(String::from))
115 })
116 .or_else(|| {
117 request
118 .headers()
119 .get(header::AUTHORIZATION)
120 .and_then(|v| v.to_str().ok())
121 .and_then(|v| v.strip_prefix("Bearer "))
122 .map(String::from)
123 })
124 .or_else(|| {
125 if request.uri().path() != "/graphql/ws" {
126 return None;
127 }
128 request.uri().query().and_then(|q| {
129 q.split('&')
130 .find_map(|pair| pair.strip_prefix("token=").map(String::from))
131 })
132 })
133}
134
135#[cfg(test)]
140mod tests {
141 use super::*;
142 use axum::body::Body;
143 use axum::http::Request as HttpRequest;
144 use axum::routing::get;
145 use koan_core::auth::Role;
146 use tower::ServiceExt as _;
147
148 async fn echo_user(axum::Extension(user): axum::Extension<AuthUser>) -> String {
150 format!("{}:{}", user.username, user.role.as_str())
151 }
152
153 async fn call(state: AuthState, req: HttpRequest<Body>) -> (StatusCode, String) {
154 let app = axum::Router::new()
155 .route("/graphql", get(echo_user))
156 .route("/graphql/ws", get(echo_user))
157 .layer(axum::middleware::from_fn_with_state(state, auth_middleware));
158 let resp = app.oneshot(req).await.unwrap();
159 let status = resp.status();
160 let bytes = axum::body::to_bytes(resp.into_body(), 64 * 1024)
161 .await
162 .unwrap();
163 (status, String::from_utf8_lossy(&bytes).into_owned())
164 }
165
166 fn keys_and_token(role: Role) -> (Vec<u8>, String) {
168 let (private_pem, public_pem) = auth::generate_keypair_pem().unwrap();
169 let token = auth::mint_access_token(private_pem.as_bytes(), 1, "alice", role, 900).unwrap();
170 (public_pem.into_bytes(), token)
171 }
172
173 fn enforcing_with(
176 public_pem: Vec<u8>,
177 key: Option<&str>,
178 username: &str,
179 role: Role,
180 ) -> (AuthState, tempfile::TempDir) {
181 let dir = tempfile::tempdir().unwrap();
182 let path = dir.path().join("koan.db");
183 let db = koan_core::db::connection::Database::open(&path).unwrap();
184 koan_core::db::queries::auth::create_user(&db.conn, username, "pw", role).unwrap();
185 let state = AuthState {
186 public_pem: Arc::new(public_pem),
187 auth_enabled: true,
188 introspection_key: key.map(|k| Arc::new(k.to_string())),
189 pool: Arc::new(Pool::new(path)),
190 };
191 (state, dir)
192 }
193
194 fn enforcing(
195 public_pem: Vec<u8>,
196 key: Option<&str>,
197 role: Role,
198 ) -> (AuthState, tempfile::TempDir) {
199 enforcing_with(public_pem, key, "alice", role)
200 }
201
202 #[tokio::test]
203 async fn auth_disabled_grants_anonymous_admin() {
204 let state = AuthState {
205 public_pem: Arc::new(Vec::new()),
206 auth_enabled: false,
207 introspection_key: None,
208 pool: Arc::new(Pool::new("/nonexistent/koan.db".into())),
209 };
210 let req = HttpRequest::get("/graphql").body(Body::empty()).unwrap();
211 let (status, body) = call(state, req).await;
212 assert_eq!(status, StatusCode::OK);
213 assert_eq!(body, "anonymous:admin");
214 }
215
216 #[tokio::test]
217 async fn missing_token_is_unauthorized() {
218 let (public_pem, _) = keys_and_token(Role::Admin);
219 let req = HttpRequest::get("/graphql").body(Body::empty()).unwrap();
220 let (state, _dir) = enforcing(public_pem, None, Role::Admin);
221 let (status, _) = call(state, req).await;
222 assert_eq!(status, StatusCode::UNAUTHORIZED);
223 }
224
225 #[tokio::test]
226 async fn bearer_token_authenticates() {
227 let (public_pem, token) = keys_and_token(Role::User);
228 let req = HttpRequest::get("/graphql")
229 .header(header::AUTHORIZATION, format!("Bearer {token}"))
230 .body(Body::empty())
231 .unwrap();
232 let (state, _dir) = enforcing(public_pem, None, Role::User);
233 let (status, body) = call(state, req).await;
234 assert_eq!(status, StatusCode::OK);
235 assert_eq!(body, "alice:user");
236 }
237
238 #[tokio::test]
239 async fn cookie_takes_precedence_over_bearer() {
240 let (public_pem, cookie_token) = keys_and_token(Role::Readonly);
241 let req = HttpRequest::get("/graphql")
242 .header(
243 header::COOKIE,
244 format!("other=1; koan_access={cookie_token}"),
245 )
246 .header(header::AUTHORIZATION, "Bearer garbage")
247 .body(Body::empty())
248 .unwrap();
249 let (state, _dir) = enforcing(public_pem, None, Role::Readonly);
250 let (status, body) = call(state, req).await;
251 assert_eq!(status, StatusCode::OK);
252 assert_eq!(body, "alice:readonly");
253 }
254
255 #[tokio::test]
256 async fn query_param_token_only_works_on_the_ws_route() {
257 let (public_pem, token) = keys_and_token(Role::Admin);
258
259 let req = HttpRequest::get(format!("/graphql?token={token}"))
260 .body(Body::empty())
261 .unwrap();
262 let (state, _dir) = enforcing(public_pem, None, Role::Admin);
263 let (status, _) = call(state.clone(), req).await;
264 assert_eq!(status, StatusCode::UNAUTHORIZED);
265
266 let req = HttpRequest::get(format!("/graphql/ws?token={token}"))
267 .body(Body::empty())
268 .unwrap();
269 let (status, body) = call(state, req).await;
270 assert_eq!(status, StatusCode::OK);
271 assert_eq!(body, "alice:admin");
272 }
273
274 #[tokio::test]
275 async fn introspection_key_bypasses_auth_only_when_it_matches() {
276 let (public_pem, _) = keys_and_token(Role::Admin);
277
278 let req = HttpRequest::get("/graphql")
279 .header("X-Introspection-Key", "sekrit")
280 .body(Body::empty())
281 .unwrap();
282 let (state, _dir) = enforcing(public_pem, Some("sekrit"), Role::Admin);
283 let (status, body) = call(state.clone(), req).await;
284 assert_eq!(status, StatusCode::OK);
285 assert_eq!(body, "anonymous:admin");
286
287 let req = HttpRequest::get("/graphql")
288 .header("X-Introspection-Key", "sekrjt")
289 .body(Body::empty())
290 .unwrap();
291 let (status, _) = call(state, req).await;
292 assert_eq!(status, StatusCode::UNAUTHORIZED);
293 }
294
295 #[tokio::test]
296 async fn tampered_token_is_rejected() {
297 let (public_pem, token) = keys_and_token(Role::Admin);
298 let req = HttpRequest::get("/graphql")
299 .header(header::AUTHORIZATION, format!("Bearer {token}x"))
300 .body(Body::empty())
301 .unwrap();
302 let (state, _dir) = enforcing(public_pem, None, Role::Admin);
303 let (status, _) = call(state, req).await;
304 assert_eq!(status, StatusCode::UNAUTHORIZED);
305 }
306
307 #[tokio::test]
308 async fn token_signed_by_another_key_is_rejected() {
309 let (_, token) = keys_and_token(Role::Admin);
310 let (other_public, _) = keys_and_token(Role::Admin);
311 let req = HttpRequest::get("/graphql")
312 .header(header::AUTHORIZATION, format!("Bearer {token}"))
313 .body(Body::empty())
314 .unwrap();
315 let (state, _dir) = enforcing(other_public, None, Role::Admin);
316 let (status, _) = call(state, req).await;
317 assert_eq!(status, StatusCode::UNAUTHORIZED);
318 }
319
320 #[tokio::test]
321 async fn the_role_is_the_accounts_now_not_the_tokens() {
322 let (public_pem, token) = keys_and_token(Role::Admin);
323 let req = || {
324 HttpRequest::get("/graphql")
325 .header(header::AUTHORIZATION, format!("Bearer {token}"))
326 .body(Body::empty())
327 .unwrap()
328 };
329 let (state, _dir) = enforcing(public_pem.clone(), None, Role::Readonly);
331 assert_eq!(
332 call(state, req()).await,
333 (StatusCode::OK, "alice:readonly".into())
334 );
335
336 let (state, dir) = enforcing(public_pem.clone(), None, Role::Admin);
338 let db = koan_core::db::connection::Database::open(&dir.path().join("koan.db")).unwrap();
339 koan_core::db::queries::auth::delete_user(&db.conn, 1).unwrap();
340 assert_eq!(call(state, req()).await.0, StatusCode::UNAUTHORIZED);
341
342 let (state, _dir) = enforcing_with(public_pem, None, "bob", Role::Admin);
344 assert_eq!(call(state, req()).await.0, StatusCode::UNAUTHORIZED);
345 }
346}