feat: add api key generation endpint + role authorization

This commit is contained in:
2026-05-07 12:32:29 +02:00
parent d7ddc087a6
commit 752373c7b7
10 changed files with 315 additions and 114 deletions
+11
View File
@@ -178,3 +178,14 @@ pub struct ChatDelta {
pub role: Option<Role>,
pub content: Option<String>,
}
#[derive(serde::Deserialize)]
pub struct CreateApiKeyRequest {
pub name: String,
pub scopes: Vec<String>,
}
#[derive(serde::Serialize)]
pub struct CreateApiKeyResponse {
pub api_key: String, // ONLY returned once
}
-58
View File
@@ -1,58 +0,0 @@
use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
use once_cell::sync::Lazy;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::env;
#[derive(Debug, Deserialize, Serialize, Clone)]
pub struct Claims {
pub sub: String,
pub preferred_username: Option<String>,
pub exp: usize,
pub iss: String,
pub aud: Option<Vec<String>>,
pub realm_access: Option<RealmAccess>,
}
#[derive(Debug, Deserialize, Serialize, Clone)]
pub struct RealmAccess {
pub roles: Vec<String>,
}
static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set"));
pub fn validate_token(token: &str, jwks: &Value) -> Result<Claims, String> {
// 1. Decode header
let header = decode_header(token).map_err(|_| "Invalid header")?;
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("Matching key not found")?;
// 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(|_| "Invalid decoding key")?;
// 4. Setup validation rules
let mut validation = Validation::new(Algorithm::RS256);
validation.set_issuer(&[ISSUER.as_str()]);
validation.validate_exp = true;
validation.validate_aud = false;
// 5. Decode & verify
let token_data = decode::<Claims>(token, &decoding_key, &validation)
.map_err(|_| "Token validation failed")?;
Ok(token_data.claims)
}
+91
View File
@@ -1,10 +1,15 @@
use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
use once_cell::sync::Lazy;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
use std::env;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
// ------ JWKS ------
#[derive(Clone)]
struct JwksCache {
jwks: Value,
@@ -55,3 +60,89 @@ pub async fn get_jwks() -> Result<Value, reqwest::Error> {
// Expired or empty → refresh
refresh_jwks().await
}
// ------ Claims ------
#[derive(Debug, Deserialize, Serialize, Clone)]
pub struct Claims {
pub sub: String,
pub preferred_username: Option<String>,
pub exp: usize,
pub iss: String,
pub aud: Option<Vec<String>>,
pub realm_access: Option<RealmAccess>,
#[serde(default)]
pub resource_access: HashMap<String, ResourceAccess>,
}
#[derive(Debug, Deserialize, Serialize, Clone)]
pub struct RealmAccess {
pub roles: Vec<String>,
}
#[derive(Debug, Deserialize, Serialize, Clone)]
pub struct ResourceAccess {
pub roles: Vec<String>,
}
impl Claims {
pub fn realm_roles(&self) -> &[String] {
self.realm_access
.as_ref()
.map_or(&[], |r| r.roles.as_slice())
}
pub fn has_realm_role(&self, role: &str) -> bool {
self.realm_roles().iter().any(|r| r == role)
}
pub fn client_roles(&self, client: &str) -> &[String] {
self.resource_access
.get(client)
.map_or(&[], |r| r.roles.as_slice())
}
pub fn has_client_role(&self, client: &str, role: &str) -> bool {
self.client_roles(client).iter().any(|r| r == role)
}
}
// ------ Claims ------
static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set"));
pub fn validate_token(token: &str, jwks: &Value) -> Result<Claims, String> {
// 1. Decode header
let header = decode_header(token).map_err(|_| "Invalid header")?;
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("Matching key not found")?;
// 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(|_| "Invalid decoding key")?;
// 4. Setup validation rules
let mut validation = Validation::new(Algorithm::RS256);
validation.set_issuer(&[ISSUER.as_str()]);
validation.validate_exp = true;
validation.validate_aud = false;
// 5. Decode & verify
let token_data = decode::<Claims>(token, &decoding_key, &validation)
.map_err(|_| "Token validation failed")?;
Ok(token_data.claims)
}
+48 -32
View File
@@ -6,17 +6,58 @@ use axum::{
};
use uuid::Uuid;
use crate::middlewares::auth::{
jwt::validate_token,
keycloak::{get_jwks, refresh_jwks},
};
use crate::middlewares::auth::keycloak::{Claims, get_jwks, refresh_jwks, validate_token};
use crate::databases::postgres::user_repository::ensure_user_exists;
use crate::state::app_state::AppState;
pub async fn require_roles(
request: Request<axum::body::Body>,
next: Next,
realm_role: Option<&'static str>,
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(
state: &AppState,
mut request: Request,
next: Next,
claims: Claims,
) -> Result<Response, StatusCode> {
dbg!(&claims);
let user_id = claims
.sub
.parse::<Uuid>()
.map_err(|_| StatusCode::UNAUTHORIZED)?;
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 async fn auth_middleware(
State(state): State<AppState>,
mut request: Request,
request: Request,
next: Next,
) -> Result<Response, StatusCode> {
tracing::debug!("Middleware hit");
@@ -40,38 +81,13 @@ pub async fn auth_middleware(
match validate_token(token, &jwks) {
Ok(claims) => {
tracing::debug!("Token valid");
// Create user in db
let user_id = claims
.sub
.parse::<Uuid>()
.map_err(|_| StatusCode::UNAUTHORIZED)?;
ensure_user_exists(&state.postgres, user_id)
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
request.extensions_mut().insert(claims);
Ok(next.run(request).await)
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) => {
let user_id = claims
.sub
.parse::<Uuid>()
.map_err(|_| StatusCode::UNAUTHORIZED)?;
ensure_user_exists(&state.postgres, user_id)
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
request.extensions_mut().insert(claims);
Ok(next.run(request).await)
}
Ok(claims) => handle_valid_claims(&state, request, next, claims).await,
Err(e) => {
tracing::error!("JWT validation failed: {:?}", e);
Err(StatusCode::UNAUTHORIZED)
-1
View File
@@ -1,4 +1,3 @@
pub mod jwt;
pub mod keycloak;
pub mod middleware;
+59
View File
@@ -0,0 +1,59 @@
use axum::{
Json,
extract::{Extension, State},
http::StatusCode,
};
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::state::app_state::AppState;
fn generate_api_key() -> String {
let mut bytes = [0u8; 32];
OsRng.fill_bytes(&mut 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(
State(state): State<AppState>,
Extension(claims): Extension<Claims>,
Json(body): Json<CreateApiKeyRequest>,
) -> Result<Json<CreateApiKeyResponse>, StatusCode> {
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)
VALUES ($1, $2, $3, $4)
"#,
key_hash,
body.name,
user_id,
&body.scopes
)
.execute(&state.postgres)
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
Ok(Json(CreateApiKeyResponse { api_key: raw_key }))
}
+9 -1
View File
@@ -1,9 +1,11 @@
pub mod apikey;
pub mod chat;
pub mod models;
use crate::docs::ApiDoc;
use crate::middlewares::auth::auth_middleware;
use crate::middlewares::auth::{auth_middleware, middleware::require_roles};
use crate::state::app_state::AppState;
use axum::{Json, Router, middleware, routing::get, routing::post};
use utoipa::OpenApi;
@@ -22,6 +24,12 @@ pub fn protected_router() -> Router<AppState> {
.route("/chat/completions", post(chat::chat_completions))
.route("/models/{model}/load", post(models::load_model))
.route("/models/{model}/unload", post(models::unload_model))
.route(
"/keys/generate",
post(apikey::create_api_key).route_layer(middleware::from_fn(|req, next| {
require_roles(req, next, None, None)
})),
)
}
pub fn router(state: AppState) -> Router<AppState> {