Compare commits

...
5 Commits
Author SHA1 Message Date
LucasDLTG 2e81132f2d fix: use nex table name
CI / Rust CI (push) Failing after 1m12s
Publish & Deploy / Build and Push to Registry (push) Failing after 1m48s
Publish & Deploy / Deploy via SSH (push) Has been skipped
2026-05-07 19:07:14 +02:00
LucasDLTG c534299c8e feat: update last time used api key 2026-05-07 16:15:27 +02:00
LucasDLTG 4b428ec32a feat: add api key verification 2026-05-07 14:13:29 +02:00
LucasDLTG 752373c7b7 feat: add api key generation endpint + role authorization 2026-05-07 12:32:29 +02:00
LucasDLTG d7ddc087a6 feat: add user in db 2026-05-06 15:22:49 +02:00
21 changed files with 1355 additions and 136 deletions
Generated
+876 -24
View File
File diff suppressed because it is too large Load Diff
+8 -4
View File
@@ -5,16 +5,16 @@ edition = "2024"
[dev-dependencies]
wiremock = "0.6"
tokio = { version = "1.52.1", features = ["macros", "rt-multi-thread"] }
tokio = { version = "1.52.2", features = ["macros", "rt-multi-thread"] }
[dependencies]
axum = "0.8.9"
utoipa = { version = "5.4.0", features = ["axum_extras"] }
utoipa = { version = "5.5.0", features = ["axum_extras"] }
tokio = { version = "1", features = ["full"] }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] }
reqwest = { version = "0.13.2", features = ["json", "stream"] }
reqwest = { version = "0.13.3", features = ["json", "stream"] }
once_cell = "1"
dotenvy = "0.15"
thiserror = "2.0.18"
@@ -22,6 +22,10 @@ tokio-stream = "0.1"
futures = "0.3"
chrono = { version = "0.4.44", features = ["serde"] }
uuid = { version = "1.23.1", features = ["v4", "serde"] }
tower-http = { version = "0.6.8", features = ["cors"] }
tower-http = { version = "0.6.10", features = ["cors"] }
tracing = "0.1.44"
tracing-subscriber = { version = "0.3", features = ["env-filter"]}
sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "postgres", "uuid", "chrono"] }
rand = "0.8"
base64 = "0.22.1"
sha2 = "0.11.0"
+4
View File
@@ -263,3 +263,7 @@ This project turns Ollama into:
👉 A local OpenAI-compatible API
👉 A controllable model runtime
👉 A foundation for a full LLM gateway
# TODO
- open api doc for bearer token
- db
+1
View File
@@ -0,0 +1 @@
pub mod postgres;
+17
View File
@@ -0,0 +1,17 @@
use sqlx::PgPool;
use uuid::Uuid;
pub async fn update_last_access(pool: &PgPool, api_key_id: Uuid) -> Result<(), sqlx::Error> {
sqlx::query!(
r#"
UPDATE auth.api_key
SET last_used_at = now()
WHERE id = $1
"#,
api_key_id
)
.execute(pool)
.await?;
Ok(())
}
+3
View File
@@ -0,0 +1,3 @@
pub mod api_key;
pub mod pool;
pub mod user_repository;
+14
View File
@@ -0,0 +1,14 @@
use sqlx::{PgPool, postgres::PgPoolOptions};
use std::time::Duration;
pub async fn create_pool(database_url: &str) -> Result<PgPool, sqlx::Error> {
PgPoolOptions::new()
.max_connections(10)
.acquire_timeout(Duration::from_secs(5))
.connect(database_url)
.await
.map_err(|err| {
tracing::error!("Postgres connection error: {:?}", err);
err
})
}
+17
View File
@@ -0,0 +1,17 @@
use sqlx::PgPool;
use uuid::Uuid;
pub async fn ensure_user_exists(pool: &PgPool, user_id: Uuid) -> Result<(), sqlx::Error> {
sqlx::query!(
r#"
INSERT INTO auth.app_user (id)
VALUES ($1)
ON CONFLICT (id) DO NOTHING
"#,
user_id
)
.execute(pool)
.await?;
Ok(())
}
+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
}
+3
View File
@@ -1,4 +1,7 @@
pub mod databases;
pub mod dto;
pub mod errors;
pub mod middlewares;
pub mod providers;
pub mod state;
pub mod utils;
+20 -4
View File
@@ -1,3 +1,4 @@
mod databases;
mod docs;
mod dto;
mod errors;
@@ -5,12 +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;
@@ -35,8 +38,16 @@ async fn main() {
init_tracing();
// DB Connection
let database_url = env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = postgres::pool::create_pool(&database_url)
.await
.expect("Fatal error");
let state = AppState {
ollama: Arc::new(OllamaProvider::new(OLLAMA_URL.as_str())),
postgres: pool,
};
let cors_origin =
@@ -45,16 +56,21 @@ async fn main() {
let cors = CorsLayer::new()
.allow_origin(cors_origin.parse::<HeaderValue>().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()
.nest("/v1", routes::v1::router())
.nest("/v1", routes::v1::router(state.clone()))
.layer(cors)
.with_state(state);
let addr = SocketAddr::from(([0, 0, 0, 0], 3001));
println!("Server running on {}", addr);
tracing::debug!("Server running on {}", addr);
axum::serve(tokio::net::TcpListener::bind(addr).await.unwrap(), app)
.await
+7
View File
@@ -0,0 +1,7 @@
use uuid::Uuid;
#[derive(Clone, Debug)]
pub struct ApiKeyClaims {
pub sub: Uuid,
pub api_key_id: Uuid,
}
-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, Default)]
pub struct KeycloakClaims {
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 KeycloakClaims {
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)
}
}
// ------ Validation ------
static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set"));
pub fn validate_token(token: &str, jwks: &Value) -> Result<KeycloakClaims, 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::<KeycloakClaims>(token, &decoding_key, &validation)
.map_err(|_| "Token validation failed")?;
Ok(token_data.claims)
}
+200 -35
View File
@@ -1,50 +1,215 @@
use axum::{extract::Request, http::StatusCode, middleware::Next, response::Response};
use crate::middlewares::auth::{
jwt::validate_token,
keycloak::{get_jwks, refresh_jwks},
use axum::{
extract::{Request, State},
http::StatusCode,
middleware::Next,
response::{IntoResponse, Response},
};
pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Response, StatusCode> {
tracing::debug!("Middleware hit");
use crate::databases::postgres::{
api_key::update_last_access, 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;
let headers = request.headers();
#[derive(Clone, Debug)]
pub enum Auth {
Jwt(KeycloakClaims),
ApiKey(ApiKeyClaims),
}
tracing::debug!("Headers extracted");
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 auth_header = headers.get("authorization").and_then(|v| v.to_str().ok());
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
}
}
tracing::debug!("Auth header: {:?}", auth_header);
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,
}
}
}
let auth_header = auth_header.ok_or(StatusCode::UNAUTHORIZED)?;
pub async fn auth_middleware(
State(state): State<AppState>,
request: Request,
next: Next,
) -> Result<Response, StatusCode> {
match try_jwt(&state, request, next).await {
Ok(response) => Ok(response),
Err((request, next)) => try_api_key(&state, request, next).await,
}
}
let token = auth_header
.strip_prefix("Bearer ")
/// 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, ak.id as key_id
FROM auth.api_key ak
JOIN auth.app_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)?;
let jwks = get_jwks().await.map_err(|_| StatusCode::UNAUTHORIZED)?;
tracing::debug!("API key valid, user_id={}", row.user_id);
match validate_token(token, &jwks) {
Ok(claims) => {
tracing::debug!("Token valid");
handle_auth(
state,
request,
next,
Auth::ApiKey(ApiKeyClaims {
sub: row.user_id,
api_key_id: row.key_id,
}),
)
.await
}
request.extensions_mut().insert(claims);
// ── 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> {
match &auth {
Auth::Jwt(_) => {
ensure_user_exists(&state.postgres, auth.user_id())
.await
.map_err(|e| {
tracing::error!("ensure_user_exists failed: {e}");
StatusCode::INTERNAL_SERVER_ERROR
})?;
}
Auth::ApiKey(key) => {
update_last_access(&state.postgres, key.api_key_id)
.await
.map_err(|e| {
tracing::error!("update_last_access 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)
}
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(e) => {
tracing::error!("JWT validation failed: {:?}", e);
Err(StatusCode::UNAUTHORIZED)
}
}
}
}
}
+1 -1
View File
@@ -1,4 +1,4 @@
pub mod jwt;
pub mod apikey;
pub mod keycloak;
pub mod middleware;
+50
View File
@@ -0,0 +1,50 @@
use axum::{
Json,
extract::{Extension, State},
http::StatusCode,
};
use base64::{Engine as _, engine::general_purpose};
use rand::RngCore;
use rand::rngs::OsRng;
use crate::dto::api::{CreateApiKeyRequest, CreateApiKeyResponse};
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];
OsRng.fill_bytes(&mut bytes);
general_purpose::URL_SAFE_NO_PAD.encode(bytes)
}
pub async fn create_api_key(
State(state): State<AppState>,
Extension(claims): Extension<Auth>,
Json(body): Json<CreateApiKeyRequest>,
) -> Result<Json<CreateApiKeyResponse>, 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);
sqlx::query!(
r#"
INSERT INTO auth.api_key (key_hash, name, created_by, scopes)
VALUES ($1, $2, $3, $4)
"#,
key_hash,
body.name,
claims.user_id(),
&body.scopes
)
.execute(&state.postgres)
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
Ok(Json(CreateApiKeyResponse { api_key: raw_key }))
}
+11 -3
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,10 +24,16 @@ 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() -> Router<AppState> {
pub fn router(state: AppState) -> Router<AppState> {
Router::new()
.merge(public_router())
.merge(protected_router().layer(middleware::from_fn(auth_middleware)))
.merge(protected_router().layer(middleware::from_fn_with_state(state, auth_middleware)))
}
+2
View File
@@ -1,7 +1,9 @@
use crate::providers::ollama::client::OllamaProvider;
use sqlx::PgPool;
use std::sync::Arc;
#[derive(Clone)]
pub struct AppState {
pub ollama: Arc<OllamaProvider>,
pub postgres: PgPool,
}
+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;