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 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
90/// Priority: `koan_access` cookie, then `Authorization: Bearer`, then `?token=`.
91///
92/// The query parameter is confined to the WebSocket route, which is the only one
93/// that cannot carry a header. A token in a URL survives in shell history, proxy
94/// logs and `Referer`.
95fn 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// ---------------------------------------------------------------------------
125// Tests
126// ---------------------------------------------------------------------------
127
128#[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    /// Echoes the `AuthUser` the middleware injected, so tests can assert on it.
138    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    /// A live keypair plus a matching token for `alice` at `role`.
156    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    /// Enforcing auth over a database whose first account is `username` at
163    /// `role`: id 1, which the tokens above name.
164    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        // Demoted since the token was minted.
319        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        // Deleted since.
326        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        // Its id now belongs to someone else.
332        let (state, _dir) = enforcing_with(public_pem, None, "bob", Role::Admin);
333        assert_eq!(call(state, req()).await.0, StatusCode::UNAUTHORIZED);
334    }
335}