diff --git a/Cargo.lock b/Cargo.lock index c61468f..623ebf2 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -168,6 +168,15 @@ dependencies = [ "generic-array", ] +[[package]] +name = "block-buffer" +version = "0.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cdd35008169921d80bc60d3d0ab416eecb028c4cd653352907921d95084790be" +dependencies = [ + "hybrid-array", +] + [[package]] name = "bumpalo" version = "3.20.2" @@ -221,14 +230,17 @@ name = "chat" version = "0.1.0" dependencies = [ "axum", + "base64", "chrono", "dotenvy", "futures", "jsonwebtoken", "once_cell", + "rand 0.8.6", "reqwest", "serde", "serde_json", + "sha2 0.11.0", "sqlx", "thiserror 2.0.18", "tokio", @@ -289,6 +301,12 @@ version = "0.9.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" +[[package]] +name = "const-oid" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c" + [[package]] name = "core-foundation" version = "0.9.4" @@ -324,6 +342,15 @@ dependencies = [ "libc", ] +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + [[package]] name = "crc" version = "3.4.0" @@ -364,6 +391,15 @@ dependencies = [ "typenum", ] +[[package]] +name = "crypto-common" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77727bb15fa921304124b128af125e7e3b968275d1b108b379190264f4423710" +dependencies = [ + "hybrid-array", +] + [[package]] name = "deadpool" version = "0.12.3" @@ -388,7 +424,7 @@ version = "0.7.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e7c1832837b905bbfb5101e07cc24c8deddf52f93225eee6ead5f4d63d53ddcb" dependencies = [ - "const-oid", + "const-oid 0.9.6", "pem-rfc7468", "zeroize", ] @@ -408,12 +444,23 @@ version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "block-buffer", - "const-oid", - "crypto-common", + "block-buffer 0.10.4", + "const-oid 0.9.6", + "crypto-common 0.1.7", "subtle", ] +[[package]] +name = "digest" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1dd6dbb5841937940781866fa1281a1ff7bd3bf827091440879f9994983d5c2" +dependencies = [ + "block-buffer 0.12.0", + "const-oid 0.10.2", + "crypto-common 0.2.1", +] + [[package]] name = "displaydoc" version = "0.2.5" @@ -764,7 +811,7 @@ version = "0.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" dependencies = [ - "digest", + "digest 0.10.7", ] [[package]] @@ -821,6 +868,15 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" +[[package]] +name = "hybrid-array" +version = "0.4.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08d46837a0ed51fe95bd3b05de33cd64a1ee88fc797477ca48446872504507c5" +dependencies = [ + "typenum", +] + [[package]] name = "hyper" version = "1.9.0" @@ -1222,7 +1278,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" dependencies = [ "cfg-if", - "digest", + "digest 0.10.7", ] [[package]] @@ -1723,8 +1779,8 @@ version = "0.9.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b8573f03f5883dcaebdfcf4725caa1ecb9c15b2ef50c43a07b816e06799bb12d" dependencies = [ - "const-oid", - "digest", + "const-oid 0.9.6", + "digest 0.10.7", "num-bigint-dig", "num-integer", "num-traits", @@ -1957,8 +2013,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" dependencies = [ "cfg-if", - "cpufeatures", - "digest", + "cpufeatures 0.2.17", + "digest 0.10.7", ] [[package]] @@ -1968,8 +2024,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", - "cpufeatures", - "digest", + "cpufeatures 0.2.17", + "digest 0.10.7", +] + +[[package]] +name = "sha2" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "446ba717509524cb3f22f17ecc096f10f4822d76ab5c0b9822c5f9c284e825f4" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "digest 0.11.3", ] [[package]] @@ -2003,7 +2070,7 @@ version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de" dependencies = [ - "digest", + "digest 0.10.7", "rand_core 0.6.4", ] @@ -2103,7 +2170,7 @@ dependencies = [ "rustls", "serde", "serde_json", - "sha2", + "sha2 0.10.9", "smallvec", "thiserror 2.0.18", "tokio", @@ -2142,7 +2209,7 @@ dependencies = [ "quote", "serde", "serde_json", - "sha2", + "sha2 0.10.9", "sqlx-core", "sqlx-mysql", "sqlx-postgres", @@ -2165,7 +2232,7 @@ dependencies = [ "bytes", "chrono", "crc", - "digest", + "digest 0.10.7", "dotenvy", "either", "futures-channel", @@ -2186,7 +2253,7 @@ dependencies = [ "rsa", "serde", "sha1", - "sha2", + "sha2 0.10.9", "smallvec", "sqlx-core", "stringprep", @@ -2225,7 +2292,7 @@ dependencies = [ "rand 0.8.6", "serde", "serde_json", - "sha2", + "sha2 0.10.9", "smallvec", "sqlx-core", "stringprep", @@ -2521,9 +2588,9 @@ dependencies = [ [[package]] name = "tower-http" -version = "0.6.9" +version = "0.6.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a28f0d049ccfaa566e14e9663d304d8577427b368cb4710a20528690287a738b" +checksum = "68d6fdd9f81c2819c9a8b0e0cd91660e7746a8e6ea2ba7c6b2b057985f6bcb51" dependencies = [ "bitflags", "bytes", diff --git a/Cargo.toml b/Cargo.toml index 1348c7b..2606fe3 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -22,7 +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.9", 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"] } \ No newline at end of file +sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "postgres", "uuid", "chrono"] } +rand = "0.8" +base64 = "0.22.1" +sha2 = "0.11.0" \ No newline at end of file diff --git a/readme.md b/readme.md index 11efdec..3373daa 100644 --- a/readme.md +++ b/readme.md @@ -263,3 +263,8 @@ 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 +- endpoint for creating token +- verify bearer token diff --git a/src/dto/api.rs b/src/dto/api.rs index cf0e199..5209edf 100644 --- a/src/dto/api.rs +++ b/src/dto/api.rs @@ -178,3 +178,14 @@ pub struct ChatDelta { pub role: Option, pub content: Option, } + +#[derive(serde::Deserialize)] +pub struct CreateApiKeyRequest { + pub name: String, + pub scopes: Vec, +} + +#[derive(serde::Serialize)] +pub struct CreateApiKeyResponse { + pub api_key: String, // ONLY returned once +} diff --git a/src/middlewares/auth/jwt.rs b/src/middlewares/auth/jwt.rs deleted file mode 100644 index 8c63d46..0000000 --- a/src/middlewares/auth/jwt.rs +++ /dev/null @@ -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, - pub exp: usize, - pub iss: String, - pub aud: Option>, - pub realm_access: Option, -} - -#[derive(Debug, Deserialize, Serialize, Clone)] -pub struct RealmAccess { - pub roles: Vec, -} - -static ISSUER: Lazy = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set")); - -pub fn validate_token(token: &str, jwks: &Value) -> Result { - // 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::(token, &decoding_key, &validation) - .map_err(|_| "Token validation failed")?; - - Ok(token_data.claims) -} diff --git a/src/middlewares/auth/keycloak.rs b/src/middlewares/auth/keycloak.rs index f8796af..b74c4b9 100644 --- a/src/middlewares/auth/keycloak.rs +++ b/src/middlewares/auth/keycloak.rs @@ -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 { // Expired or empty → refresh refresh_jwks().await } + +// ------ Claims ------ + +#[derive(Debug, Deserialize, Serialize, Clone)] +pub struct Claims { + pub sub: String, + pub preferred_username: Option, + pub exp: usize, + pub iss: String, + pub aud: Option>, + pub realm_access: Option, + #[serde(default)] + pub resource_access: HashMap, +} + +#[derive(Debug, Deserialize, Serialize, Clone)] +pub struct RealmAccess { + pub roles: Vec, +} + +#[derive(Debug, Deserialize, Serialize, Clone)] +pub struct ResourceAccess { + pub roles: Vec, +} + +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 = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set")); + +pub fn validate_token(token: &str, jwks: &Value) -> Result { + // 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::(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 92ea504..c2f9fbe 100644 --- a/src/middlewares/auth/middleware.rs +++ b/src/middlewares/auth/middleware.rs @@ -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, + 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) +} + +async fn handle_valid_claims( + state: &AppState, + mut request: Request, + next: Next, + claims: Claims, +) -> Result { + dbg!(&claims); + + let user_id = claims + .sub + .parse::() + .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, - mut request: Request, + request: Request, next: Next, ) -> Result { 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::() - .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::() - .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) diff --git a/src/middlewares/auth/mod.rs b/src/middlewares/auth/mod.rs index f09dc4f..52031ed 100644 --- a/src/middlewares/auth/mod.rs +++ b/src/middlewares/auth/mod.rs @@ -1,4 +1,3 @@ -pub mod jwt; pub mod keycloak; pub mod middleware; diff --git a/src/routes/v1/apikey.rs b/src/routes/v1/apikey.rs new file mode 100644 index 0000000..35fbb0b --- /dev/null +++ b/src/routes/v1/apikey.rs @@ -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, + Extension(claims): Extension, + Json(body): Json, +) -> Result, 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 })) +} diff --git a/src/routes/v1/mod.rs b/src/routes/v1/mod.rs index f83d6eb..3acb327 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -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 { .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 {