diff --git a/Cargo.lock b/Cargo.lock index 87e4843fa9..43f12ef434 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1202,6 +1202,7 @@ dependencies = [ "serde_json", "serde_with", "size", + "strum", "tokio", "tokio-stream", "tracing", diff --git a/Cargo.toml b/Cargo.toml index 87ef7f1e92..f8ff89bf82 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -31,6 +31,7 @@ hulyrs = { git = "https://github.com/hcengineering/hulyrs.git", features = [ ] } secrecy = "0.10.3" tokio-stream = "0.1" +strum = { version = "0.27.2", features = ["derive"] } [[bin]] name = "hulypulse" diff --git a/src/config.rs b/src/config.rs index 181a96df5e..e9ca6186ce 100644 --- a/src/config.rs +++ b/src/config.rs @@ -13,15 +13,14 @@ // limitations under the License. // +use std::{path::Path, sync::LazyLock}; + use secrecy::SecretString; use serde::Deserialize; use serde_with::formats::CommaSeparator; use serde_with::{StringWithSeparator, serde_as}; - use url::Url; -use std::{path::Path, sync::LazyLock}; - use config::FileFormat; #[derive(Deserialize, Debug, PartialEq)] @@ -31,6 +30,13 @@ pub enum RedisMode { Direct, } +#[derive(Deserialize, Debug, PartialEq, strum::Display)] +#[serde(rename_all = "lowercase")] +pub enum BackendType { + Memory, + Redis, +} + #[serde_as] #[derive(Deserialize, Debug)] pub struct Config { @@ -48,8 +54,8 @@ pub struct Config { pub max_ttl: usize, pub max_size: Option, - pub memory_mode: Option, - pub no_authorization: Option, + pub backend: BackendType, + pub no_authorization: bool, } pub static CONFIG: LazyLock = LazyLock::new(|| { diff --git a/src/config/default.toml b/src/config/default.toml index e3a6e64fd1..2e7980fb5b 100644 --- a/src/config/default.toml +++ b/src/config/default.toml @@ -9,9 +9,8 @@ redis_mode = "direct" redis_service = "mymaster" max_ttl = 3600 +backend = "redis" +no_authorization = false # optional settings - -# memory_mode = true -# no_authorization = true -# max_size = 100 \ No newline at end of file +# max_size = 100 diff --git a/src/handlers_ws.rs b/src/handlers_ws.rs index dde3f6c04e..2ffe771b55 100644 --- a/src/handlers_ws.rs +++ b/src/handlers_ws.rs @@ -507,22 +507,22 @@ pub async fn handler( db: web::Data, hub_state: web::Data>>, ) -> Result { - let claims = if CONFIG.no_authorization == Some(true) { - None - } else { + let claims = if !CONFIG.no_authorization { Some( req.extensions() .get::() .expect("Missing claims") .to_owned(), ) + } else { + None }; let session = WsSession { db: db.get_ref().clone(), hub_state: hub_state.get_ref().clone(), id: new_session_id(), - claims: claims, + claims, }; ws::start(session, &req, payload) diff --git a/src/main.rs b/src/main.rs index 4b4a06a348..4de348e724 100644 --- a/src/main.rs +++ b/src/main.rs @@ -67,19 +67,18 @@ async fn extract_claims( token: Option, } - if CONFIG.no_authorization == Some(true) { - return next.call(request).await; + if !CONFIG.no_authorization { + let query = request.extract::>().await?.into_inner(); + + let claims = if let Some(token) = query.token { + Claims::from_token(token, CONFIG.token_secret.expose_secret()).unwrap() + } else { + request.extract_claims(&CONFIG.token_secret)? + }; + + request.extensions_mut().insert(claims); } - let query = request.extract::>().await?.into_inner(); - - let claims = if let Some(token) = query.token { - Claims::from_token(token, CONFIG.token_secret.expose_secret()).unwrap() - } else { - request.extract_claims(&CONFIG.token_secret)? - }; - - request.extensions_mut().insert(claims); next.call(request).await } @@ -87,22 +86,22 @@ async fn check_workspace( mut request: ServiceRequest, next: Next, ) -> Result, Error> { - if CONFIG.no_authorization.unwrap_or(false) { - return next.call(request).await; - } + if !CONFIG.no_authorization { + let workspace = Uuid::parse_str(&request.extract::>().await?); + let claims = request.extensions().get::().cloned().unwrap(); - let workspace = Uuid::parse_str(&request.extract::>().await?); - let claims = request.extensions().get::().cloned().unwrap(); - - if claims.is_system() || Ok(claims.workspace.clone()) == workspace.clone().map(Some) { - next.call(request).await + if claims.is_system() || Ok(claims.workspace.clone()) == workspace.clone().map(Some) { + next.call(request).await + } else { + warn!( + expected = ?claims.workspace, + actual = ?workspace, + "Unauthorized request, workspace mismatch" + ); + Err(actix_web::error::ErrorUnauthorized("Unauthorized").into()) + } } else { - warn!( - expected = ?claims.workspace, - actual = ?workspace, - "Unauthorized request, workspace mismatch" - ); - Err(actix_web::error::ErrorUnauthorized("Unauthorized").into()) + next.call(request).await } } @@ -115,17 +114,21 @@ async fn main() -> anyhow::Result<()> { // starting HubService let hub_state = Arc::new(RwLock::new(HubState::default())); - let db_backend = if CONFIG.memory_mode == Some(true) { - let memory = MemoryBackend::new(); - memory.spawn_ticker(hub_state.clone()); - tracing::info!("Memory mode enabled"); - Db::new_memory(memory, hub_state.clone()) - } else { - let redis_client = redis::client().await?; - let redis_connection = redis_client.get_multiplexed_async_connection().await?; - tokio::spawn(crate::redis::receiver(redis_client, hub_state.clone())); - tracing::info!("Redis mode enabled"); - Db::new_redis(redis_connection, hub_state.clone()) + let db_backend = match CONFIG.backend { + config::BackendType::Memory => { + let memory = MemoryBackend::new(); + memory.spawn_ticker(hub_state.clone()); + tracing::info!("Memory mode enabled"); + Db::new_memory(memory, hub_state.clone()) + } + + config::BackendType::Redis => { + let redis_client = redis::client().await?; + let redis_connection = redis_client.get_multiplexed_async_connection().await?; + tokio::spawn(crate::redis::receiver(redis_client, hub_state.clone())); + tracing::info!("Redis mode enabled"); + Db::new_redis(redis_connection, hub_state.clone()) + } }; let socket = std::net::SocketAddr::new(CONFIG.bind_host.as_str().parse()?, CONFIG.bind_port); @@ -180,11 +183,7 @@ async fn main() -> anyhow::Result<()> { let count = hub_state.read().await.count(); Ok::<_, actix_web::Error>(HttpResponse::Ok().json(json!({ "memory_info": info, - "db_mode": if CONFIG.memory_mode == Some(true) { - "memory" - } else { - "redis" - }, + "backend": CONFIG.backend.to_string().to_lowercase(), "websockets": count, "status": "OK", }))) diff --git a/src/workspace_owner.rs b/src/workspace_owner.rs index a294dfb8b9..b0611a8cc0 100644 --- a/src/workspace_owner.rs +++ b/src/workspace_owner.rs @@ -24,7 +24,7 @@ pub fn check_workspace_core(claims_opt: Option, key: &str) -> Result<(), return Err("Invalid key: deprecated symbols"); } - if CONFIG.no_authorization == Some(true) { + if CONFIG.no_authorization { return Ok(()); }