feat: add user in db
This commit is contained in:
Generated
+809
-24
File diff suppressed because it is too large
Load Diff
+6
-5
@@ -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"] }
|
||||||
@@ -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 dto;
|
||||||
pub mod errors;
|
pub mod errors;
|
||||||
pub mod middlewares;
|
pub mod middlewares;
|
||||||
pub mod providers;
|
pub mod providers;
|
||||||
|
pub mod state;
|
||||||
|
|||||||
+12
-2
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user