feat: add user in db
This commit is contained in:
@@ -0,0 +1 @@
|
||||
pub mod postgres;
|
||||
@@ -0,0 +1,2 @@
|
||||
pub mod pool;
|
||||
pub mod user_repository;
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -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(())
|
||||
}
|
||||
@@ -1,4 +1,6 @@
|
||||
pub mod databases;
|
||||
pub mod dto;
|
||||
pub mod errors;
|
||||
pub mod middlewares;
|
||||
pub mod providers;
|
||||
pub mod state;
|
||||
|
||||
+12
-2
@@ -1,3 +1,4 @@
|
||||
mod databases;
|
||||
mod docs;
|
||||
mod dto;
|
||||
mod errors;
|
||||
@@ -6,6 +7,7 @@ mod providers;
|
||||
mod routes;
|
||||
mod state;
|
||||
|
||||
use crate::databases::postgres;
|
||||
use crate::providers::ollama::client::OllamaProvider;
|
||||
use crate::state::app_state::AppState;
|
||||
|
||||
@@ -35,8 +37,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 =
|
||||
@@ -49,12 +59,12 @@ async fn main() {
|
||||
.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
|
||||
|
||||
@@ -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::{
|
||||
jwt::validate_token,
|
||||
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");
|
||||
|
||||
let headers = request.headers();
|
||||
@@ -28,6 +41,16 @@ pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Respons
|
||||
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)
|
||||
@@ -37,6 +60,15 @@ pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Respons
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -24,8 +24,8 @@ pub fn protected_router() -> Router<AppState> {
|
||||
.route("/models/{model}/unload", post(models::unload_model))
|
||||
}
|
||||
|
||||
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)))
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user