feat: add api key verification

This commit is contained in:
2026-05-07 14:13:29 +02:00
parent 752373c7b7
commit 4b428ec32a
9 changed files with 209 additions and 92 deletions
+1
View File
@@ -4,3 +4,4 @@ pub mod errors;
pub mod middlewares; pub mod middlewares;
pub mod providers; pub mod providers;
pub mod state; pub mod state;
pub mod utils;
+8 -2
View File
@@ -6,13 +6,14 @@ mod middlewares;
mod providers; mod providers;
mod routes; mod routes;
mod state; mod state;
mod utils;
use crate::databases::postgres; use crate::databases::postgres;
use crate::providers::ollama::client::OllamaProvider; use crate::providers::ollama::client::OllamaProvider;
use crate::state::app_state::AppState; use crate::state::app_state::AppState;
use axum::Router; use axum::Router;
use axum::http::{HeaderValue, Method, header}; use axum::http::{HeaderName, HeaderValue, Method, header};
use once_cell::sync::Lazy; use once_cell::sync::Lazy;
use std::env; use std::env;
use std::net::SocketAddr; use std::net::SocketAddr;
@@ -55,7 +56,12 @@ async fn main() {
let cors = CorsLayer::new() let cors = CorsLayer::new()
.allow_origin(cors_origin.parse::<HeaderValue>().unwrap()) .allow_origin(cors_origin.parse::<HeaderValue>().unwrap())
.allow_methods([Method::GET, Method::POST, Method::PUT, Method::DELETE]) .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); .allow_credentials(true);
let app = Router::new() let app = Router::new()
+6
View File
@@ -0,0 +1,6 @@
use uuid::Uuid;
#[derive(Clone, Debug)]
pub struct ApiKeyClaims {
pub sub: Uuid,
}
+6 -6
View File
@@ -63,8 +63,8 @@ pub async fn get_jwks() -> Result<Value, reqwest::Error> {
// ------ Claims ------ // ------ Claims ------
#[derive(Debug, Deserialize, Serialize, Clone)] #[derive(Debug, Deserialize, Serialize, Clone, Default)]
pub struct Claims { pub struct KeycloakClaims {
pub sub: String, pub sub: String,
pub preferred_username: Option<String>, pub preferred_username: Option<String>,
pub exp: usize, pub exp: usize,
@@ -85,7 +85,7 @@ pub struct ResourceAccess {
pub roles: Vec<String>, pub roles: Vec<String>,
} }
impl Claims { impl KeycloakClaims {
pub fn realm_roles(&self) -> &[String] { pub fn realm_roles(&self) -> &[String] {
self.realm_access self.realm_access
.as_ref() .as_ref()
@@ -107,11 +107,11 @@ impl Claims {
} }
} }
// ------ Claims ------ // ------ Validation ------
static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set")); static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set"));
pub fn validate_token(token: &str, jwks: &Value) -> Result<Claims, String> { pub fn validate_token(token: &str, jwks: &Value) -> Result<KeycloakClaims, String> {
// 1. Decode header // 1. Decode header
let header = decode_header(token).map_err(|_| "Invalid header")?; let header = decode_header(token).map_err(|_| "Invalid header")?;
@@ -141,7 +141,7 @@ pub fn validate_token(token: &str, jwks: &Value) -> Result<Claims, String> {
validation.validate_aud = false; validation.validate_aud = false;
// 5. Decode & verify // 5. Decode & verify
let token_data = decode::<Claims>(token, &decoding_key, &validation) let token_data = decode::<KeycloakClaims>(token, &decoding_key, &validation)
.map_err(|_| "Token validation failed")?; .map_err(|_| "Token validation failed")?;
Ok(token_data.claims) Ok(token_data.claims)
+176 -76
View File
@@ -2,57 +2,43 @@ use axum::{
extract::{Request, State}, extract::{Request, State},
http::StatusCode, http::StatusCode,
middleware::Next, 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::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::state::app_state::AppState;
use crate::utils::crypto::hash_key;
use uuid::Uuid;
pub async fn require_roles( #[derive(Clone, Debug)]
request: Request<axum::body::Body>, pub enum Auth {
next: Next, Jwt(KeycloakClaims),
realm_role: Option<&'static str>, ApiKey(ApiKeyClaims),
client_role: Option<&'static str>,
) -> Result<Response, axum::http::StatusCode> {
let claims = request
.extensions()
.get::<Claims>()
.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)
} }
async fn handle_valid_claims( impl Auth {
state: &AppState, pub fn user_id(&self) -> Uuid {
mut request: Request, match self {
next: Next, Auth::Jwt(c) => c.sub.parse().expect("sub is a valid UUID"),
claims: Claims, Auth::ApiKey(c) => c.sub,
) -> Result<Response, StatusCode> { }
dbg!(&claims); }
let user_id = claims pub fn has_realm_role(&self, role: &str) -> bool {
.sub match self {
.parse::<Uuid>() Auth::Jwt(c) => c.has_realm_role(role),
.map_err(|_| StatusCode::UNAUTHORIZED)?; Auth::ApiKey(_) => false, // API keys carry no roles
}
}
ensure_user_exists(&state.postgres, user_id) pub fn has_client_role(&self, client: &str, role: &str) -> bool {
.await match self {
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; Auth::Jwt(c) => c.has_client_role(client, role),
Auth::ApiKey(_) => false,
request.extensions_mut().insert(claims); }
}
Ok(next.run(request).await)
} }
pub async fn auth_middleware( pub async fn auth_middleware(
@@ -60,39 +46,153 @@ pub async fn auth_middleware(
request: Request, request: Request,
next: Next, next: Next,
) -> Result<Response, StatusCode> { ) -> Result<Response, StatusCode> {
tracing::debug!("Middleware hit"); match try_jwt(&state, request, next).await {
Ok(response) => Ok(response),
let headers = request.headers(); Err((request, next)) => try_api_key(&state, request, next).await,
tracing::debug!("Headers extracted");
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
}
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,
Err(e) => {
tracing::error!("JWT validation failed: {:?}", e);
Err(StatusCode::UNAUTHORIZED)
}
}
}
} }
} }
/// 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<Response, (Request, Next)> {
let token = request
.headers()
.get("authorization")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "))
.map(str::to_owned);
let Some(token) = token else {
// No Authorization header at all → let API key branch try
return Err((request, next));
};
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());
}
};
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!("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<Response, StatusCode> {
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<Response, StatusCode> {
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<Response, StatusCode> {
let auth = request
.extensions()
.get::<Auth>()
.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)
}
+1
View File
@@ -1,3 +1,4 @@
pub mod apikey;
pub mod keycloak; pub mod keycloak;
pub mod middleware; pub mod middleware;
+8 -17
View File
@@ -6,12 +6,11 @@ use axum::{
use base64::{Engine as _, engine::general_purpose}; use base64::{Engine as _, engine::general_purpose};
use rand::RngCore; use rand::RngCore;
use rand::rngs::OsRng; use rand::rngs::OsRng;
use sha2::{Digest, Sha256};
use uuid::Uuid;
use crate::dto::api::{CreateApiKeyRequest, CreateApiKeyResponse}; 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::state::app_state::AppState;
use crate::utils::crypto::hash_key;
fn generate_api_key() -> String { fn generate_api_key() -> String {
let mut bytes = [0u8; 32]; let mut bytes = [0u8; 32];
@@ -19,28 +18,20 @@ fn generate_api_key() -> String {
general_purpose::URL_SAFE_NO_PAD.encode(bytes) 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( pub async fn create_api_key(
State(state): State<AppState>, State(state): State<AppState>,
Extension(claims): Extension<Claims>, Extension(claims): Extension<Auth>,
Json(body): Json<CreateApiKeyRequest>, Json(body): Json<CreateApiKeyRequest>,
) -> Result<Json<CreateApiKeyResponse>, StatusCode> { ) -> Result<Json<CreateApiKeyResponse>, StatusCode> {
if matches!(claims, Auth::ApiKey(_)) {
return Err(StatusCode::FORBIDDEN);
}
let raw_key = generate_api_key(); let raw_key = generate_api_key();
let key_hash = hash_key(&raw_key); let key_hash = hash_key(&raw_key);
dbg!(&claims); dbg!(&claims);
let user_id = Uuid::parse_str(&claims.sub).map_err(|_| StatusCode::UNAUTHORIZED)?;
sqlx::query!( sqlx::query!(
r#" r#"
INSERT INTO auth.api_key (key_hash, name, created_by, scopes) INSERT INTO auth.api_key (key_hash, name, created_by, scopes)
@@ -48,7 +39,7 @@ pub async fn create_api_key(
"#, "#,
key_hash, key_hash,
body.name, body.name,
user_id, claims.user_id(),
&body.scopes &body.scopes
) )
.execute(&state.postgres) .execute(&state.postgres)
+11
View File
@@ -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()
}
+1
View File
@@ -0,0 +1 @@
pub mod crypto;