From 05bc471d01bdc9787bf6cebacb11f52f99642d36 Mon Sep 17 00:00:00 2001 From: LucasX Ubuntu Date: Thu, 9 Apr 2026 15:31:55 +0200 Subject: [PATCH] feat: add auth middleware --- .env | 1 - .env.example | 2 ++ .gitignore | 1 + Cargo.lock | 12 ++++++-- Cargo.toml | 2 +- readme.md | 3 ++ src/auth/jwt.rs | 47 +++++++++++++++++++++--------- src/auth/keycloak.rs | 32 +++++++++++--------- src/auth/middleware.rs | 54 ++++++++++++++++++++++++++++++++++ src/auth/mod.rs | 2 +- src/main.rs | 66 +++++++++++++++++++++++++++++++++--------- 11 files changed, 176 insertions(+), 46 deletions(-) delete mode 100644 .env create mode 100644 .env.example create mode 100644 readme.md create mode 100644 src/auth/middleware.rs diff --git a/.env b/.env deleted file mode 100644 index d45b4c7..0000000 --- a/.env +++ /dev/null @@ -1 +0,0 @@ -JWKS_URL=https://auth.iceberg.black/realms/iceberg/protocol/openid-connect/certs \ No newline at end of file diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..4f9bdad --- /dev/null +++ b/.env.example @@ -0,0 +1,2 @@ +JWKS_URL=https://auth.iceberg.black/realms/iceberg/protocol/openid-connect/certs +ISSUER=https://auth.iceberg.black/realms/iceberg \ No newline at end of file diff --git a/.gitignore b/.gitignore index ea8c4bf..0b745e2 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,2 @@ /target +.env \ No newline at end of file diff --git a/Cargo.lock b/Cargo.lock index 0523ce4..c4d9d26 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -21,6 +21,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a054912289d18629dc78375ba2c3726a3afe3ff71b4edba9dedfca0e3446d1fc" dependencies = [ "aws-lc-sys", + "untrusted 0.7.1", "zeroize", ] @@ -691,6 +692,7 @@ version = "10.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0529410abe238729a60b108898784df8984c87f6054c9c4fcacc47e4803c1ce1" dependencies = [ + "aws-lc-rs", "base64", "getrandom 0.2.17", "js-sys", @@ -1055,7 +1057,7 @@ dependencies = [ "cfg-if", "getrandom 0.2.17", "libc", - "untrusted", + "untrusted 0.9.0", "windows-sys 0.52.0", ] @@ -1137,7 +1139,7 @@ dependencies = [ "aws-lc-rs", "ring", "rustls-pki-types", - "untrusted", + "untrusted 0.9.0", ] [[package]] @@ -1614,6 +1616,12 @@ version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "untrusted" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a156c684c91ea7d62626509bce3cb4e1d9ed5c4d978f7b4352658f96a4c26b4a" + [[package]] name = "untrusted" version = "0.9.0" diff --git a/Cargo.toml b/Cargo.toml index c6c0619..743da72 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -8,7 +8,7 @@ axum = "0.8.8" tokio = { version = "1", features = ["full"] } serde = { version = "1", features = ["derive"] } serde_json = "1" -jsonwebtoken = "10.3.0" +jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] } reqwest = { version = "0.13.2", features = ["json"] } once_cell = "1" dotenvy = "0.15" \ No newline at end of file diff --git a/readme.md b/readme.md new file mode 100644 index 0000000..4c03b82 --- /dev/null +++ b/readme.md @@ -0,0 +1,3 @@ +# TODO + +- Race condition on jwks token refresh \ No newline at end of file diff --git a/src/auth/jwt.rs b/src/auth/jwt.rs index c851186..21eb32c 100644 --- a/src/auth/jwt.rs +++ b/src/auth/jwt.rs @@ -1,9 +1,10 @@ use jsonwebtoken::{decode, decode_header, Algorithm, DecodingKey, Validation}; use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::env; +use once_cell::sync::Lazy; -static ISSUER: &str = "https://auth.iceberg.black/realms/iceberg"; - -#[derive(Debug, Deserialize, Serialize)] +#[derive(Debug, Deserialize, Serialize, Clone)] pub struct Claims { pub sub: String, pub preferred_username: Option, @@ -13,34 +14,52 @@ pub struct Claims { pub realm_access: Option, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Debug, Deserialize, Serialize, Clone)] pub struct RealmAccess { pub roles: Vec, } -pub fn validate_token(token: &str, jwks: &serde_json::Value) -> Result { - let header = decode_header(token).map_err(|_| ())?; +static ISSUER: Lazy = Lazy::new(|| { + env::var("ISSUER").expect("ISSUER not set") +}); - let kid = header.kid.ok_or(())?; +pub fn validate_token(token: &str, jwks: &Value) -> Result { + // 1. Decode header + let header = decode_header(token).map_err(|_| "Invalid header")?; - let keys = jwks["keys"].as_array().ok_or(())?; + let kid = header.kid.ok_or("Missing kid")?; + + // 2. Find matching key + let keys = jwks["keys"] + .as_array() + .ok_or("Invalid JWKS")?; let key = keys .iter() .find(|k| k["kid"] == kid) - .ok_or(())?; + .ok_or("Matching key not found")?; - let n = key["n"].as_str().ok_or(())?; - let e = key["e"].as_str().ok_or(())?; + // 3. Extract RSA components + let n = key["n"].as_str().ok_or("Missing n")?; + let e = key["e"].as_str().ok_or("Missing e")?; let decoding_key = - DecodingKey::from_rsa_components(n, e).map_err(|_| ())?; + DecodingKey::from_rsa_components(n, e) + .map_err(|_| "Invalid decoding key")?; + // 4. Setup validation rules let mut validation = Validation::new(Algorithm::RS256); - validation.set_issuer(&[ISSUER]); + validation.set_issuer(&[ISSUER.as_str()]); + + // Optional but recommended: + validation.validate_exp = true; + validation.validate_aud = false; // depends on your Keycloak config + + // 5. Decode & verify let token_data = - decode::(token, &decoding_key, &validation).map_err(|_| ())?; + decode::(token, &decoding_key, &validation) + .map_err(|_| "Token validation failed")?; Ok(token_data.claims) } \ No newline at end of file diff --git a/src/auth/keycloak.rs b/src/auth/keycloak.rs index b208767..56846f6 100644 --- a/src/auth/keycloak.rs +++ b/src/auth/keycloak.rs @@ -14,10 +14,12 @@ struct JwksCache { static JWK_CACHE: Lazy>>> = Lazy::new(|| Arc::new(RwLock::new(None))); -async fn fetch_jwks() -> Result { - let url = env::var("JWKS_URL").expect("JWKS_URL not set"); +static JWKS_URL: Lazy = Lazy::new(|| { + env::var("JWKS_URL").expect("JWKS_URL not set") +}); - let jwks = reqwest::get(url) +async fn fetch_jwks() -> Result { + let jwks = reqwest::get(JWKS_URL.as_str()) .await? .json::() .await?; @@ -25,6 +27,19 @@ async fn fetch_jwks() -> Result { Ok(jwks) } +pub async fn refresh_jwks() -> Result { + let jwks = fetch_jwks().await?; + + let mut write = JWK_CACHE.write().await; + + *write = Some(JwksCache { + jwks: jwks.clone(), + last_fetched: Instant::now(), + }); + + Ok(jwks) +} + pub async fn get_jwks() -> Result { let ttl = Duration::from_secs(3600); // 1 hour @@ -40,14 +55,5 @@ pub async fn get_jwks() -> Result { } // Expired or empty → refresh - let jwks = fetch_jwks().await?; - - let mut write = JWK_CACHE.write().await; - - *write = Some(JwksCache { - jwks: jwks.clone(), - last_fetched: Instant::now(), - }); - - Ok(jwks) + refresh_jwks().await } \ No newline at end of file diff --git a/src/auth/middleware.rs b/src/auth/middleware.rs new file mode 100644 index 0000000..e9072c3 --- /dev/null +++ b/src/auth/middleware.rs @@ -0,0 +1,54 @@ +use axum::{ + extract::Request, + http::{StatusCode}, + middleware::Next, + response::Response, +}; + +use crate::auth::{jwt::validate_token, keycloak::{get_jwks, refresh_jwks}}; + +pub async fn auth_middleware( + mut request: Request, + next: Next, +) -> Result { + dbg!("Middleware hit"); + + let headers = request.headers(); + + dbg!("Headers extracted"); + + let auth_header = headers + .get("authorization") + .and_then(|v| v.to_str().ok()); + + dbg!("Auth header: {:?}", auth_header); + + let auth_header = auth_header.ok_or(StatusCode::UNAUTHORIZED)?; + + let token = auth_header + .strip_prefix("Bearer ") + .ok_or(StatusCode::UNAUTHORIZED)?; + + let jwks = get_jwks().await.map_err(|_| StatusCode::UNAUTHORIZED)?; + + match validate_token(token, &jwks) { + Ok(claims) => { + dbg!("Token valid"); + + request.extensions_mut().insert(claims); + + Ok(next.run(request).await) + } + Err(_) => { + let jwks = refresh_jwks().await.map_err(|_| StatusCode::UNAUTHORIZED)?; + + match validate_token(token, &jwks) { + Ok(claims) => { + request.extensions_mut().insert(claims); + Ok(next.run(request).await) + } + Err(_) => Err(StatusCode::UNAUTHORIZED), + } + } + } +} \ No newline at end of file diff --git a/src/auth/mod.rs b/src/auth/mod.rs index f00d87b..c260d26 100644 --- a/src/auth/mod.rs +++ b/src/auth/mod.rs @@ -1,3 +1,3 @@ pub mod jwt; pub mod keycloak; -// pub mod middleware; \ No newline at end of file +pub mod middleware; \ No newline at end of file diff --git a/src/main.rs b/src/main.rs index e164079..ef137e8 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,34 +1,72 @@ mod auth; -use crate::auth::{keycloak, jwt}; +use crate::auth::middleware::auth_middleware; +use crate::auth::jwt::Claims; +use dotenvy::dotenv; +use axum::extract::Extension; use axum::{ - extract::Request, - http::{HeaderMap, StatusCode}, routing::get, Router, + middleware, }; -use dotenvy::dotenv; -use std::env; +use std::net::SocketAddr; +pub async fn protected_route( + Extension(claims): Extension, +) -> String { + format!( + "Hello {}, your user id is {}", + claims + .preferred_username + .unwrap_or("unknown".to_string()), + claims.sub + ) +} +pub async fn public_route() -> &'static str { + println!("Public route hit"); + "Public endpoint: no authentication required" +} +pub fn app() -> Router { + let public_routes = Router::new() + .route("/", get(public_route)); + + let protected_routes = Router::new() + .route("/protected", get(protected_route)) + .layer(middleware::from_fn(auth_middleware)); + + Router::new() + .merge(public_routes) + .merge(protected_routes) +} #[tokio::main] async fn main() { dotenv().ok(); - let jwks = keycloak::get_jwks().await; +// let jwks = keycloak::get_jwks() +// .await +// .expect("Failed to fetch JWKS"); - println!("{:?}", jwks); +// // println!("{:?}", jwks); - // let app = Router::new().route("/protected", get(protected_route)) - // .route("/", get(dumb)); +// let token = "eyJhbGciOiJSUzI1NiIsInR5cCIgOiAiSldUIiwia2lkIiA6ICJublpLek04TkZHVmpWbGFPRXZpMUtFSTVHQWRwaGlsYjh3RHRLeG5JOENZIn0.eyJleHAiOjE3NzU3Mzg2MDUsImlhdCI6MTc3NTczODMwNSwianRpIjoiNjk2OTY4NzQtZWMwNi00NGFkLTg0MDYtYmY3YWM4MjI5MjkxIiwiaXNzIjoiaHR0cHM6Ly9hdXRoLmljZWJlcmcuYmxhY2svcmVhbG1zL2ljZWJlcmciLCJzdWIiOiJmZGRiN2FjZC1kMmE5LTRmMTctOWIxNi1kZjVlN2EzNDI4YjciLCJ0eXAiOiJCZWFyZXIiLCJhenAiOiJjaGF0LWFwaSIsInNjb3BlIjoiIiwiY2xpZW50SG9zdCI6Ijg2LjIxMi44NC4xOTEiLCJjbGllbnRBZGRyZXNzIjoiODYuMjEyLjg0LjE5MSIsImNsaWVudF9pZCI6ImNoYXQtYXBpIn0.BS7ohLWiMDxAUz_Q-Qi2UoLYbNn8AUrYeSWeO-602SQ-AYBW3gfYxXOSeRgWyn4VfObpVfK7QfqQBUxorXxi1JVld-4fGXL8NXQNyq5Ip_JHNG1p02Z39Pe9MmC9MXOwA_GQF2PIkLIdOJ_W_guXVhl2ptEWPPSiXM5Z5CNg8lyOiKPI0g2JWV6FBRG-HMXzqnxAb1j8wGUpC9JzGwAU3sjWBGhT1AAovs-XLmm5hZEPxI-Ia3SmUnF-QjFMmebPVxLdxL7OszzVEhKipsZRiwQxjY6eJhJFFa8uycBigHPSzu_HqqkK6AjNlyExvR0EGvl9zUWdOfMPDiVX2Sg92g"; +// match jwt::validate_token(token, &jwks) { +// Ok(claims) => { +// println!("Valid token for user: {:?}", claims); +// } +// Err(err) => { +// println!("Invalid token: {}", err); +// } +// } - // let addr = SocketAddr::from(([127, 0, 0, 1], 3000)); - // println!("Server running on {}", addr); + let app = app(); + let addr = SocketAddr::from(([127, 0, 0, 1], 3000)); + println!("Server running on {}", addr); - // axum::serve(tokio::net::TcpListener::bind(addr).await.unwrap(), app) - // .await - // .unwrap(); + axum::serve(tokio::net::TcpListener::bind(addr).await.unwrap(), app) + .await + .unwrap(); } \ No newline at end of file