Skip to main content

koan_server/auth/
middleware.rs

1//! Axum middleware for JWT authentication.
2//!
3//! Extracts a token from the `koan_access` cookie or `Authorization: Bearer`,
4//! validates it, and injects `AuthUser` into request extensions.
5
6use 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/// Shared state for the auth middleware.
20#[derive(Clone)]
21pub struct AuthState {
22    /// Ed25519 public key PEM for JWT verification.
23    pub public_pem: Arc<Vec<u8>>,
24    /// Whether auth is enforced.
25    pub auth_enabled: bool,
26    /// Process-scoped introspection key. Bypasses auth when matched.
27    /// Generated randomly on server start, dies with the process.
28    pub introspection_key: Option<Arc<String>>,
29    /// Where a token's account is looked up; see `super::current_user`.
30    pub pool: Arc<Pool>,
31}
32
33/// Axum middleware: validate JWT and inject `AuthUser`.
34///
35/// When `auth_enabled = false`, injects anonymous admin and passes through.
36/// When `auth_enabled = true`, requires a valid token.
37pub 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    // Check for introspection key (playground bypass).
48    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
101/// Priority: `koan_access` cookie, then `Authorization: Bearer`, then `?token=`.
102///
103/// The query parameter is confined to the WebSocket route, which is the only one
104/// that cannot carry a header. A token in a URL survives in shell history, proxy
105/// logs and `Referer`.
106fn 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// ---------------------------------------------------------------------------
136// Tests
137// ---------------------------------------------------------------------------
138
139#[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    /// Echoes the `AuthUser` the middleware injected, so tests can assert on it.
149    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    /// A live keypair plus a matching token for `alice` at `role`.
167    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    /// Enforcing auth over a database whose first account is `username` at
174    /// `role`: id 1, which the tokens above name.
175    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        // Demoted since the token was minted.
330        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        // Deleted since.
337        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        // Its id now belongs to someone else.
343        let (state, _dir) = enforcing_with(public_pem, None, "bob", Role::Admin);
344        assert_eq!(call(state, req()).await.0, StatusCode::UNAUTHORIZED);
345    }
346}