feat: add user in db

This commit is contained in:
2026-05-06 15:22:49 +02:00
parent 7f77bfef4e
commit d7ddc087a6
11 changed files with 901 additions and 35 deletions
Generated
+809 -24
View File
File diff suppressed because it is too large Load Diff
+6 -5
View File
@@ -5,16 +5,16 @@ edition = "2024"
[dev-dependencies] [dev-dependencies]
wiremock = "0.6" wiremock = "0.6"
tokio = { version = "1.52.1", features = ["macros", "rt-multi-thread"] } tokio = { version = "1.52.2", features = ["macros", "rt-multi-thread"] }
[dependencies] [dependencies]
axum = "0.8.9" 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"] } tokio = { version = "1", features = ["full"] }
serde = { version = "1", features = ["derive"] } serde = { version = "1", features = ["derive"] }
serde_json = "1" serde_json = "1"
jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] } 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" once_cell = "1"
dotenvy = "0.15" dotenvy = "0.15"
thiserror = "2.0.18" thiserror = "2.0.18"
@@ -22,6 +22,7 @@ tokio-stream = "0.1"
futures = "0.3" futures = "0.3"
chrono = { version = "0.4.44", features = ["serde"] } chrono = { version = "0.4.44", features = ["serde"] }
uuid = { version = "1.23.1", features = ["v4", "serde"] } uuid = { version = "1.23.1", features = ["v4", "serde"] }
tower-http = { version = "0.6.8", features = ["cors"] } tower-http = { version = "0.6.9", features = ["cors"] }
tracing = "0.1.44" tracing = "0.1.44"
tracing-subscriber = { version = "0.3", features = ["env-filter"]} tracing-subscriber = { version = "0.3", features = ["env-filter"]}
sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "postgres", "uuid", "chrono"] }
+1
View File
@@ -0,0 +1 @@
pub mod postgres;
+2
View File
@@ -0,0 +1,2 @@
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.user (id)
VALUES ($1)
ON CONFLICT (id) DO NOTHING
"#,
user_id
)
.execute(pool)
.await?;
Ok(())
}
+2
View File
@@ -1,4 +1,6 @@
pub mod databases;
pub mod dto; pub mod dto;
pub mod errors; pub mod errors;
pub mod middlewares; pub mod middlewares;
pub mod providers; pub mod providers;
pub mod state;
+12 -2
View File
@@ -1,3 +1,4 @@
mod databases;
mod docs; mod docs;
mod dto; mod dto;
mod errors; mod errors;
@@ -6,6 +7,7 @@ mod providers;
mod routes; mod routes;
mod state; mod state;
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;
@@ -35,8 +37,16 @@ async fn main() {
init_tracing(); 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 { let state = AppState {
ollama: Arc::new(OllamaProvider::new(OLLAMA_URL.as_str())), ollama: Arc::new(OllamaProvider::new(OLLAMA_URL.as_str())),
postgres: pool,
}; };
let cors_origin = let cors_origin =
@@ -49,12 +59,12 @@ async fn main() {
.allow_credentials(true); .allow_credentials(true);
let app = Router::new() let app = Router::new()
.nest("/v1", routes::v1::router()) .nest("/v1", routes::v1::router(state.clone()))
.layer(cors) .layer(cors)
.with_state(state); .with_state(state);
let addr = SocketAddr::from(([0, 0, 0, 0], 3001)); 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) axum::serve(tokio::net::TcpListener::bind(addr).await.unwrap(), app)
.await .await
+34 -2
View File
@@ -1,11 +1,24 @@
use axum::{extract::Request, http::StatusCode, middleware::Next, response::Response}; use axum::{
extract::{Request, State},
http::StatusCode,
middleware::Next,
response::Response,
};
use uuid::Uuid;
use crate::middlewares::auth::{ use crate::middlewares::auth::{
jwt::validate_token, jwt::validate_token,
keycloak::{get_jwks, refresh_jwks}, keycloak::{get_jwks, refresh_jwks},
}; };
pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Response, StatusCode> { use crate::databases::postgres::user_repository::ensure_user_exists;
use crate::state::app_state::AppState;
pub async fn auth_middleware(
State(state): State<AppState>,
mut request: Request,
next: Next,
) -> Result<Response, StatusCode> {
tracing::debug!("Middleware hit"); tracing::debug!("Middleware hit");
let headers = request.headers(); let headers = request.headers();
@@ -28,6 +41,16 @@ pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Respons
Ok(claims) => { Ok(claims) => {
tracing::debug!("Token valid"); 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); request.extensions_mut().insert(claims);
Ok(next.run(request).await) Ok(next.run(request).await)
@@ -37,6 +60,15 @@ pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Respons
match validate_token(token, &jwks) { match validate_token(token, &jwks) {
Ok(claims) => { 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); request.extensions_mut().insert(claims);
Ok(next.run(request).await) Ok(next.run(request).await)
} }
+2 -2
View File
@@ -24,8 +24,8 @@ pub fn protected_router() -> Router<AppState> {
.route("/models/{model}/unload", post(models::unload_model)) .route("/models/{model}/unload", post(models::unload_model))
} }
pub fn router() -> Router<AppState> { pub fn router(state: AppState) -> Router<AppState> {
Router::new() Router::new()
.merge(public_router()) .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 crate::providers::ollama::client::OllamaProvider;
use sqlx::PgPool;
use std::sync::Arc; use std::sync::Arc;
#[derive(Clone)] #[derive(Clone)]
pub struct AppState { pub struct AppState {
pub ollama: Arc<OllamaProvider>, pub ollama: Arc<OllamaProvider>,
pub postgres: PgPool,
} }