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 user = match auth::validate_access_token(&state.public_pem, &token) {
73 Ok(claims) => super::current_user(&state.pool, claims).await,
74 Err(_) => None,
75 };
76 match user {
77 Some(user) => {
78 request.extensions_mut().insert(user);
79 next.run(request).await
80 }
81 None => (
82 StatusCode::UNAUTHORIZED,
83 [("WWW-Authenticate", "Bearer")],
84 "invalid or expired token",
85 )
86 .into_response(),
87 }
88}
89
90fn extract_token(request: &Request) -> Option<String> {
96 request
97 .headers()
98 .get(header::COOKIE)
99 .and_then(|v| v.to_str().ok())
100 .and_then(|cookies| {
101 cookies
102 .split(';')
103 .find_map(|c| c.trim().strip_prefix("koan_access=").map(String::from))
104 })
105 .or_else(|| {
106 request
107 .headers()
108 .get(header::AUTHORIZATION)
109 .and_then(|v| v.to_str().ok())
110 .and_then(|v| v.strip_prefix("Bearer "))
111 .map(String::from)
112 })
113 .or_else(|| {
114 if request.uri().path() != "/graphql/ws" {
115 return None;
116 }
117 request.uri().query().and_then(|q| {
118 q.split('&')
119 .find_map(|pair| pair.strip_prefix("token=").map(String::from))
120 })
121 })
122}
123
124#[cfg(test)]
129mod tests {
130 use super::*;
131 use axum::body::Body;
132 use axum::http::Request as HttpRequest;
133 use axum::routing::get;
134 use koan_core::auth::Role;
135 use tower::ServiceExt as _;
136
137 async fn echo_user(axum::Extension(user): axum::Extension<AuthUser>) -> String {
139 format!("{}:{}", user.username, user.role.as_str())
140 }
141
142 async fn call(state: AuthState, req: HttpRequest<Body>) -> (StatusCode, String) {
143 let app = axum::Router::new()
144 .route("/graphql", get(echo_user))
145 .route("/graphql/ws", get(echo_user))
146 .layer(axum::middleware::from_fn_with_state(state, auth_middleware));
147 let resp = app.oneshot(req).await.unwrap();
148 let status = resp.status();
149 let bytes = axum::body::to_bytes(resp.into_body(), 64 * 1024)
150 .await
151 .unwrap();
152 (status, String::from_utf8_lossy(&bytes).into_owned())
153 }
154
155 fn keys_and_token(role: Role) -> (Vec<u8>, String) {
157 let (private_pem, public_pem) = auth::generate_keypair_pem().unwrap();
158 let token = auth::mint_access_token(private_pem.as_bytes(), 1, "alice", role, 900).unwrap();
159 (public_pem.into_bytes(), token)
160 }
161
162 fn enforcing_with(
165 public_pem: Vec<u8>,
166 key: Option<&str>,
167 username: &str,
168 role: Role,
169 ) -> (AuthState, tempfile::TempDir) {
170 let dir = tempfile::tempdir().unwrap();
171 let path = dir.path().join("koan.db");
172 let db = koan_core::db::connection::Database::open(&path).unwrap();
173 koan_core::db::queries::auth::create_user(&db.conn, username, "pw", role).unwrap();
174 let state = AuthState {
175 public_pem: Arc::new(public_pem),
176 auth_enabled: true,
177 introspection_key: key.map(|k| Arc::new(k.to_string())),
178 pool: Arc::new(Pool::new(path)),
179 };
180 (state, dir)
181 }
182
183 fn enforcing(
184 public_pem: Vec<u8>,
185 key: Option<&str>,
186 role: Role,
187 ) -> (AuthState, tempfile::TempDir) {
188 enforcing_with(public_pem, key, "alice", role)
189 }
190
191 #[tokio::test]
192 async fn auth_disabled_grants_anonymous_admin() {
193 let state = AuthState {
194 public_pem: Arc::new(Vec::new()),
195 auth_enabled: false,
196 introspection_key: None,
197 pool: Arc::new(Pool::new("/nonexistent/koan.db".into())),
198 };
199 let req = HttpRequest::get("/graphql").body(Body::empty()).unwrap();
200 let (status, body) = call(state, req).await;
201 assert_eq!(status, StatusCode::OK);
202 assert_eq!(body, "anonymous:admin");
203 }
204
205 #[tokio::test]
206 async fn missing_token_is_unauthorized() {
207 let (public_pem, _) = keys_and_token(Role::Admin);
208 let req = HttpRequest::get("/graphql").body(Body::empty()).unwrap();
209 let (state, _dir) = enforcing(public_pem, None, Role::Admin);
210 let (status, _) = call(state, req).await;
211 assert_eq!(status, StatusCode::UNAUTHORIZED);
212 }
213
214 #[tokio::test]
215 async fn bearer_token_authenticates() {
216 let (public_pem, token) = keys_and_token(Role::User);
217 let req = HttpRequest::get("/graphql")
218 .header(header::AUTHORIZATION, format!("Bearer {token}"))
219 .body(Body::empty())
220 .unwrap();
221 let (state, _dir) = enforcing(public_pem, None, Role::User);
222 let (status, body) = call(state, req).await;
223 assert_eq!(status, StatusCode::OK);
224 assert_eq!(body, "alice:user");
225 }
226
227 #[tokio::test]
228 async fn cookie_takes_precedence_over_bearer() {
229 let (public_pem, cookie_token) = keys_and_token(Role::Readonly);
230 let req = HttpRequest::get("/graphql")
231 .header(
232 header::COOKIE,
233 format!("other=1; koan_access={cookie_token}"),
234 )
235 .header(header::AUTHORIZATION, "Bearer garbage")
236 .body(Body::empty())
237 .unwrap();
238 let (state, _dir) = enforcing(public_pem, None, Role::Readonly);
239 let (status, body) = call(state, req).await;
240 assert_eq!(status, StatusCode::OK);
241 assert_eq!(body, "alice:readonly");
242 }
243
244 #[tokio::test]
245 async fn query_param_token_only_works_on_the_ws_route() {
246 let (public_pem, token) = keys_and_token(Role::Admin);
247
248 let req = HttpRequest::get(format!("/graphql?token={token}"))
249 .body(Body::empty())
250 .unwrap();
251 let (state, _dir) = enforcing(public_pem, None, Role::Admin);
252 let (status, _) = call(state.clone(), req).await;
253 assert_eq!(status, StatusCode::UNAUTHORIZED);
254
255 let req = HttpRequest::get(format!("/graphql/ws?token={token}"))
256 .body(Body::empty())
257 .unwrap();
258 let (status, body) = call(state, req).await;
259 assert_eq!(status, StatusCode::OK);
260 assert_eq!(body, "alice:admin");
261 }
262
263 #[tokio::test]
264 async fn introspection_key_bypasses_auth_only_when_it_matches() {
265 let (public_pem, _) = keys_and_token(Role::Admin);
266
267 let req = HttpRequest::get("/graphql")
268 .header("X-Introspection-Key", "sekrit")
269 .body(Body::empty())
270 .unwrap();
271 let (state, _dir) = enforcing(public_pem, Some("sekrit"), Role::Admin);
272 let (status, body) = call(state.clone(), req).await;
273 assert_eq!(status, StatusCode::OK);
274 assert_eq!(body, "anonymous:admin");
275
276 let req = HttpRequest::get("/graphql")
277 .header("X-Introspection-Key", "sekrjt")
278 .body(Body::empty())
279 .unwrap();
280 let (status, _) = call(state, req).await;
281 assert_eq!(status, StatusCode::UNAUTHORIZED);
282 }
283
284 #[tokio::test]
285 async fn tampered_token_is_rejected() {
286 let (public_pem, token) = keys_and_token(Role::Admin);
287 let req = HttpRequest::get("/graphql")
288 .header(header::AUTHORIZATION, format!("Bearer {token}x"))
289 .body(Body::empty())
290 .unwrap();
291 let (state, _dir) = enforcing(public_pem, None, Role::Admin);
292 let (status, _) = call(state, req).await;
293 assert_eq!(status, StatusCode::UNAUTHORIZED);
294 }
295
296 #[tokio::test]
297 async fn token_signed_by_another_key_is_rejected() {
298 let (_, token) = keys_and_token(Role::Admin);
299 let (other_public, _) = keys_and_token(Role::Admin);
300 let req = HttpRequest::get("/graphql")
301 .header(header::AUTHORIZATION, format!("Bearer {token}"))
302 .body(Body::empty())
303 .unwrap();
304 let (state, _dir) = enforcing(other_public, None, Role::Admin);
305 let (status, _) = call(state, req).await;
306 assert_eq!(status, StatusCode::UNAUTHORIZED);
307 }
308
309 #[tokio::test]
310 async fn the_role_is_the_accounts_now_not_the_tokens() {
311 let (public_pem, token) = keys_and_token(Role::Admin);
312 let req = || {
313 HttpRequest::get("/graphql")
314 .header(header::AUTHORIZATION, format!("Bearer {token}"))
315 .body(Body::empty())
316 .unwrap()
317 };
318 let (state, _dir) = enforcing(public_pem.clone(), None, Role::Readonly);
320 assert_eq!(
321 call(state, req()).await,
322 (StatusCode::OK, "alice:readonly".into())
323 );
324
325 let (state, dir) = enforcing(public_pem.clone(), None, Role::Admin);
327 let db = koan_core::db::connection::Database::open(&dir.path().join("koan.db")).unwrap();
328 koan_core::db::queries::auth::delete_user(&db.conn, 1).unwrap();
329 assert_eq!(call(state, req()).await.0, StatusCode::UNAUTHORIZED);
330
331 let (state, _dir) = enforcing_with(public_pem, None, "bob", Role::Admin);
333 assert_eq!(call(state, req()).await.0, StatusCode::UNAUTHORIZED);
334 }
335}