feat: add api key verification
This commit is contained in:
@@ -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
@@ -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()
|
||||||
|
|||||||
@@ -0,0 +1,6 @@
|
|||||||
|
use uuid::Uuid;
|
||||||
|
|
||||||
|
#[derive(Clone, Debug)]
|
||||||
|
pub struct ApiKeyClaims {
|
||||||
|
pub sub: Uuid,
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
|||||||
@@ -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),
|
||||||
|
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<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);
|
||||||
|
|
||||||
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());
|
let jwks = match get_jwks().await {
|
||||||
|
Ok(j) => j,
|
||||||
tracing::debug!("Auth header: {:?}", auth_header);
|
Err(e) => {
|
||||||
|
tracing::error!("Failed to fetch JWKS: {e}");
|
||||||
let auth_header = auth_header.ok_or(StatusCode::UNAUTHORIZED)?;
|
// Token was present but we can't validate → hard 500
|
||||||
|
return Ok(StatusCode::INTERNAL_SERVER_ERROR.into_response());
|
||||||
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) {
|
let claims = match validate_token(&token, &jwks) {
|
||||||
Ok(claims) => handle_valid_claims(&state, request, next, claims).await,
|
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) => {
|
Err(e) => {
|
||||||
tracing::error!("JWT validation failed: {:?}", e);
|
tracing::error!("Failed to refresh JWKS: {e}");
|
||||||
Err(StatusCode::UNAUTHORIZED)
|
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,3 +1,4 @@
|
|||||||
|
pub mod apikey;
|
||||||
pub mod keycloak;
|
pub mod keycloak;
|
||||||
pub mod middleware;
|
pub mod middleware;
|
||||||
|
|
||||||
|
|||||||
+8
-17
@@ -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)
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
pub mod crypto;
|
||||||
Reference in New Issue
Block a user