From 4b428ec32aa35187d9078446df05ea237d823055 Mon Sep 17 00:00:00 2001 From: LucasDLTG Date: Thu, 7 May 2026 14:13:29 +0200 Subject: [PATCH] feat: add api key verification --- src/lib.rs | 1 + src/main.rs | 10 +- src/middlewares/auth/apikey.rs | 6 + src/middlewares/auth/keycloak.rs | 12 +- src/middlewares/auth/middleware.rs | 234 ++++++++++++++++++++--------- src/middlewares/auth/mod.rs | 1 + src/routes/v1/apikey.rs | 25 +-- src/utils/crypto.rs | 11 ++ src/utils/mod.rs | 1 + 9 files changed, 209 insertions(+), 92 deletions(-) create mode 100644 src/middlewares/auth/apikey.rs create mode 100644 src/utils/crypto.rs create mode 100644 src/utils/mod.rs diff --git a/src/lib.rs b/src/lib.rs index aaf1bf2..5abe37d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -4,3 +4,4 @@ pub mod errors; pub mod middlewares; pub mod providers; pub mod state; +pub mod utils; diff --git a/src/main.rs b/src/main.rs index aecbf25..afc5f1a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -6,13 +6,14 @@ mod middlewares; mod providers; mod routes; mod state; +mod utils; use crate::databases::postgres; use crate::providers::ollama::client::OllamaProvider; use crate::state::app_state::AppState; use axum::Router; -use axum::http::{HeaderValue, Method, header}; +use axum::http::{HeaderName, HeaderValue, Method, header}; use once_cell::sync::Lazy; use std::env; use std::net::SocketAddr; @@ -55,7 +56,12 @@ async fn main() { let cors = CorsLayer::new() .allow_origin(cors_origin.parse::().unwrap()) .allow_methods([Method::GET, Method::POST, Method::PUT, Method::DELETE]) - .allow_headers([header::CONTENT_TYPE, header::AUTHORIZATION, header::ACCEPT]) + .allow_headers([ + header::CONTENT_TYPE, + header::AUTHORIZATION, + header::ACCEPT, + HeaderName::from_static("x-api-key"), + ]) .allow_credentials(true); let app = Router::new() diff --git a/src/middlewares/auth/apikey.rs b/src/middlewares/auth/apikey.rs new file mode 100644 index 0000000..0c66985 --- /dev/null +++ b/src/middlewares/auth/apikey.rs @@ -0,0 +1,6 @@ +use uuid::Uuid; + +#[derive(Clone, Debug)] +pub struct ApiKeyClaims { + pub sub: Uuid, +} diff --git a/src/middlewares/auth/keycloak.rs b/src/middlewares/auth/keycloak.rs index b74c4b9..157f297 100644 --- a/src/middlewares/auth/keycloak.rs +++ b/src/middlewares/auth/keycloak.rs @@ -63,8 +63,8 @@ pub async fn get_jwks() -> Result { // ------ Claims ------ -#[derive(Debug, Deserialize, Serialize, Clone)] -pub struct Claims { +#[derive(Debug, Deserialize, Serialize, Clone, Default)] +pub struct KeycloakClaims { pub sub: String, pub preferred_username: Option, pub exp: usize, @@ -85,7 +85,7 @@ pub struct ResourceAccess { pub roles: Vec, } -impl Claims { +impl KeycloakClaims { pub fn realm_roles(&self) -> &[String] { self.realm_access .as_ref() @@ -107,11 +107,11 @@ impl Claims { } } -// ------ Claims ------ +// ------ Validation ------ static ISSUER: Lazy = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set")); -pub fn validate_token(token: &str, jwks: &Value) -> Result { +pub fn validate_token(token: &str, jwks: &Value) -> Result { // 1. Decode header let header = decode_header(token).map_err(|_| "Invalid header")?; @@ -141,7 +141,7 @@ pub fn validate_token(token: &str, jwks: &Value) -> Result { validation.validate_aud = false; // 5. Decode & verify - let token_data = decode::(token, &decoding_key, &validation) + let token_data = decode::(token, &decoding_key, &validation) .map_err(|_| "Token validation failed")?; Ok(token_data.claims) diff --git a/src/middlewares/auth/middleware.rs b/src/middlewares/auth/middleware.rs index c2f9fbe..7b211e6 100644 --- a/src/middlewares/auth/middleware.rs +++ b/src/middlewares/auth/middleware.rs @@ -2,57 +2,43 @@ use axum::{ extract::{Request, State}, http::StatusCode, middleware::Next, - response::Response, + response::{IntoResponse, Response}, }; -use uuid::Uuid; - -use crate::middlewares::auth::keycloak::{Claims, get_jwks, refresh_jwks, validate_token}; use crate::databases::postgres::user_repository::ensure_user_exists; +use crate::middlewares::auth::apikey::ApiKeyClaims; +use crate::middlewares::auth::keycloak::{KeycloakClaims, get_jwks, refresh_jwks, validate_token}; use crate::state::app_state::AppState; +use crate::utils::crypto::hash_key; +use uuid::Uuid; -pub async fn require_roles( - request: Request, - next: Next, - realm_role: Option<&'static str>, - client_role: Option<&'static str>, -) -> Result { - let claims = request - .extensions() - .get::() - .ok_or(axum::http::StatusCode::UNAUTHORIZED)?; - - if realm_role.is_some_and(|role| !claims.has_realm_role(role)) { - return Err(axum::http::StatusCode::FORBIDDEN); - } - - if client_role.is_some_and(|role| !claims.has_client_role("chat-api", role)) { - return Err(axum::http::StatusCode::FORBIDDEN); - } - - Ok(next.run(request).await) +#[derive(Clone, Debug)] +pub enum Auth { + Jwt(KeycloakClaims), + ApiKey(ApiKeyClaims), } -async fn handle_valid_claims( - state: &AppState, - mut request: Request, - next: Next, - claims: Claims, -) -> Result { - dbg!(&claims); +impl Auth { + pub fn user_id(&self) -> Uuid { + match self { + Auth::Jwt(c) => c.sub.parse().expect("sub is a valid UUID"), + Auth::ApiKey(c) => c.sub, + } + } - let user_id = claims - .sub - .parse::() - .map_err(|_| StatusCode::UNAUTHORIZED)?; + pub fn has_realm_role(&self, role: &str) -> bool { + match self { + Auth::Jwt(c) => c.has_realm_role(role), + Auth::ApiKey(_) => false, // API keys carry no roles + } + } - ensure_user_exists(&state.postgres, user_id) - .await - .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; - - request.extensions_mut().insert(claims); - - Ok(next.run(request).await) + pub fn has_client_role(&self, client: &str, role: &str) -> bool { + match self { + Auth::Jwt(c) => c.has_client_role(client, role), + Auth::ApiKey(_) => false, + } + } } pub async fn auth_middleware( @@ -60,39 +46,153 @@ pub async fn auth_middleware( request: Request, next: Next, ) -> Result { - tracing::debug!("Middleware hit"); + match try_jwt(&state, request, next).await { + Ok(response) => Ok(response), + Err((request, next)) => try_api_key(&state, request, next).await, + } +} - let headers = request.headers(); +/// Returns Ok(Response) if JWT was valid and request handled. +/// Returns Err((request, next)) if no JWT was present (caller should try next method). +/// Returns a 401/500 response directly if JWT was present but invalid. +async fn try_jwt( + state: &AppState, + request: Request, + next: Next, +) -> Result { + let token = request + .headers() + .get("authorization") + .and_then(|v| v.to_str().ok()) + .and_then(|v| v.strip_prefix("Bearer ")) + .map(str::to_owned); - tracing::debug!("Headers extracted"); + let Some(token) = token else { + // No Authorization header at all → let API key branch try + return Err((request, next)); + }; - let auth_header = headers.get("authorization").and_then(|v| v.to_str().ok()); - - tracing::debug!("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) => { - tracing::debug!("Token valid"); - handle_valid_claims(&state, request, next, claims).await + let jwks = match get_jwks().await { + Ok(j) => j, + Err(e) => { + tracing::error!("Failed to fetch JWKS: {e}"); + // Token was present but we can't validate → hard 500 + return Ok(StatusCode::INTERNAL_SERVER_ERROR.into_response()); } - Err(_) => { - let jwks = refresh_jwks().await.map_err(|_| StatusCode::UNAUTHORIZED)?; + }; - match validate_token(token, &jwks) { - Ok(claims) => handle_valid_claims(&state, request, next, claims).await, + let claims = match validate_token(&token, &jwks) { + Ok(c) => c, + Err(_) => { + // Try refreshing JWKS once + match refresh_jwks().await { + Ok(fresh_jwks) => match validate_token(&token, &fresh_jwks) { + Ok(c) => c, + Err(_) => { + tracing::warn!("JWT validation failed after JWKS refresh"); + return Ok(StatusCode::UNAUTHORIZED.into_response()); + } + }, Err(e) => { - tracing::error!("JWT validation failed: {:?}", e); - Err(StatusCode::UNAUTHORIZED) + tracing::error!("Failed to refresh JWKS: {e}"); + return Ok(StatusCode::INTERNAL_SERVER_ERROR.into_response()); } } } + }; + + tracing::debug!("JWT valid, sub={}", claims.sub); + handle_auth(state, request, next, Auth::Jwt(claims)) + .await + .map_err(|_| unreachable!()) +} + +/// Returns Ok(Response) if API key was valid. +/// Returns Err(StatusCode) otherwise (UNAUTHORIZED or INTERNAL_SERVER_ERROR). +async fn try_api_key( + state: &AppState, + request: Request, + next: Next, +) -> Result { + let key = request + .headers() + .get("x-api-key") + .and_then(|v| v.to_str().ok()) + .ok_or(StatusCode::UNAUTHORIZED)? + .to_owned(); + + let key_hash = hash_key(&key); + + let row = sqlx::query!( + r#" + SELECT u.id AS user_id + FROM auth.api_key ak + JOIN auth.user u ON u.id = ak.created_by + WHERE ak.key_hash = $1 + AND ak.revoked_at IS NULL + "#, + key_hash + ) + .fetch_optional(&state.postgres) + .await + .map_err(|e| { + tracing::error!("DB error during API key lookup: {e}"); + StatusCode::INTERNAL_SERVER_ERROR + })? + .ok_or(StatusCode::UNAUTHORIZED)?; + + tracing::debug!("API key valid, user_id={}", row.user_id); + + handle_auth( + state, + request, + next, + Auth::ApiKey(ApiKeyClaims { sub: row.user_id }), + ) + .await +} + +// ── Shared post-auth logic ──────────────────────────────────────────────────── + +/// Ensures the user exists in the DB, inserts `Auth` into extensions, runs the next handler. +async fn handle_auth( + state: &AppState, + mut request: Request, + next: Next, + auth: Auth, +) -> Result { + ensure_user_exists(&state.postgres, auth.user_id()) + .await + .map_err(|e| { + tracing::error!("ensure_user_exists failed: {e}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + + request.extensions_mut().insert(auth); + Ok(next.run(request).await) +} + +// ── Role guard ─────────────────────────────────────────────────────────────── + +/// Layer-level middleware that checks roles *after* `auth_middleware` has run. +pub async fn require_roles( + request: Request, + next: Next, + realm_role: Option<&'static str>, + client_role: Option<&'static str>, +) -> Result { + let auth = request + .extensions() + .get::() + .ok_or(StatusCode::UNAUTHORIZED)?; + + if realm_role.is_some_and(|role| !auth.has_realm_role(role)) { + return Err(StatusCode::FORBIDDEN); } + + if client_role.is_some_and(|role| !auth.has_client_role("chat-api", role)) { + return Err(StatusCode::FORBIDDEN); + } + + Ok(next.run(request).await) } diff --git a/src/middlewares/auth/mod.rs b/src/middlewares/auth/mod.rs index 52031ed..a8b3d11 100644 --- a/src/middlewares/auth/mod.rs +++ b/src/middlewares/auth/mod.rs @@ -1,3 +1,4 @@ +pub mod apikey; pub mod keycloak; pub mod middleware; diff --git a/src/routes/v1/apikey.rs b/src/routes/v1/apikey.rs index 35fbb0b..b12c28b 100644 --- a/src/routes/v1/apikey.rs +++ b/src/routes/v1/apikey.rs @@ -6,12 +6,11 @@ use axum::{ use base64::{Engine as _, engine::general_purpose}; use rand::RngCore; use rand::rngs::OsRng; -use sha2::{Digest, Sha256}; -use uuid::Uuid; use crate::dto::api::{CreateApiKeyRequest, CreateApiKeyResponse}; -use crate::middlewares::auth::keycloak::Claims; +use crate::middlewares::auth::middleware::Auth; use crate::state::app_state::AppState; +use crate::utils::crypto::hash_key; fn generate_api_key() -> String { let mut bytes = [0u8; 32]; @@ -19,28 +18,20 @@ fn generate_api_key() -> String { general_purpose::URL_SAFE_NO_PAD.encode(bytes) } -fn hash_key(key: &str) -> String { - let mut hasher = Sha256::new(); - hasher.update(key.as_bytes()); - hasher - .finalize() - .iter() - .map(|b| format!("{:02x}", b)) - .collect() -} - pub async fn create_api_key( State(state): State, - Extension(claims): Extension, + Extension(claims): Extension, Json(body): Json, ) -> Result, StatusCode> { + if matches!(claims, Auth::ApiKey(_)) { + return Err(StatusCode::FORBIDDEN); + } + let raw_key = generate_api_key(); let key_hash = hash_key(&raw_key); dbg!(&claims); - let user_id = Uuid::parse_str(&claims.sub).map_err(|_| StatusCode::UNAUTHORIZED)?; - sqlx::query!( r#" INSERT INTO auth.api_key (key_hash, name, created_by, scopes) @@ -48,7 +39,7 @@ pub async fn create_api_key( "#, key_hash, body.name, - user_id, + claims.user_id(), &body.scopes ) .execute(&state.postgres) diff --git a/src/utils/crypto.rs b/src/utils/crypto.rs new file mode 100644 index 0000000..c9c0df2 --- /dev/null +++ b/src/utils/crypto.rs @@ -0,0 +1,11 @@ +use sha2::{Digest, Sha256}; + +pub fn hash_key(key: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(key.as_bytes()); + hasher + .finalize() + .iter() + .map(|b| format!("{:02x}", b)) + .collect() +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs new file mode 100644 index 0000000..274f0ed --- /dev/null +++ b/src/utils/mod.rs @@ -0,0 +1 @@ +pub mod crypto;