diff --git a/.sqlx/query-64e2cd0e9e56065007eeac8cf038b68cb9f0fd6d689f3eb34de23f5859797ca2.json b/.sqlx/query-64e2cd0e9e56065007eeac8cf038b68cb9f0fd6d689f3eb34de23f5859797ca2.json new file mode 100644 index 0000000000..41428929e1 --- /dev/null +++ b/.sqlx/query-64e2cd0e9e56065007eeac8cf038b68cb9f0fd6d689f3eb34de23f5859797ca2.json @@ -0,0 +1,37 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT group_id, client_traffic_policy \"client_traffic_policy: ClientTrafficPolicy\" FROM group_client_traffic_policy ORDER BY group_id", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "group_id", + "type_info": "Int8" + }, + { + "ordinal": 1, + "name": "client_traffic_policy: ClientTrafficPolicy", + "type_info": { + "Custom": { + "name": "client_traffic_policy", + "kind": { + "Enum": [ + "none", + "disable_all_traffic", + "force_all_traffic" + ] + } + } + } + } + ], + "parameters": { + "Left": [] + }, + "nullable": [ + false, + false + ] + }, + "hash": "64e2cd0e9e56065007eeac8cf038b68cb9f0fd6d689f3eb34de23f5859797ca2" +} diff --git a/.sqlx/query-93ddc32a014dc4e232842416fa6b1b277daa23b14073edba46e4f9bde969bb91.json b/.sqlx/query-93ddc32a014dc4e232842416fa6b1b277daa23b14073edba46e4f9bde969bb91.json new file mode 100644 index 0000000000..e3fba057ac --- /dev/null +++ b/.sqlx/query-93ddc32a014dc4e232842416fa6b1b277daa23b14073edba46e4f9bde969bb91.json @@ -0,0 +1,12 @@ +{ + "db_name": "PostgreSQL", + "query": "DELETE FROM group_client_traffic_policy", + "describe": { + "columns": [], + "parameters": { + "Left": [] + }, + "nullable": [] + }, + "hash": "93ddc32a014dc4e232842416fa6b1b277daa23b14073edba46e4f9bde969bb91" +} diff --git a/.sqlx/query-a7fe7af3323984a260dbc3fab701781eebf0ffb1687bc97d0f42e0a97f420d63.json b/.sqlx/query-a7fe7af3323984a260dbc3fab701781eebf0ffb1687bc97d0f42e0a97f420d63.json new file mode 100644 index 0000000000..3d90726fee --- /dev/null +++ b/.sqlx/query-a7fe7af3323984a260dbc3fab701781eebf0ffb1687bc97d0f42e0a97f420d63.json @@ -0,0 +1,51 @@ +{ + "db_name": "PostgreSQL", + "query": "INSERT INTO group_client_traffic_policy (group_id, client_traffic_policy) VALUES ($1, $2) ON CONFLICT (group_id) DO UPDATE SET client_traffic_policy = EXCLUDED.client_traffic_policy RETURNING group_id, client_traffic_policy \"client_traffic_policy: ClientTrafficPolicy\"", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "group_id", + "type_info": "Int8" + }, + { + "ordinal": 1, + "name": "client_traffic_policy: ClientTrafficPolicy", + "type_info": { + "Custom": { + "name": "client_traffic_policy", + "kind": { + "Enum": [ + "none", + "disable_all_traffic", + "force_all_traffic" + ] + } + } + } + } + ], + "parameters": { + "Left": [ + "Int8", + { + "Custom": { + "name": "client_traffic_policy", + "kind": { + "Enum": [ + "none", + "disable_all_traffic", + "force_all_traffic" + ] + } + } + } + ] + }, + "nullable": [ + false, + false + ] + }, + "hash": "a7fe7af3323984a260dbc3fab701781eebf0ffb1687bc97d0f42e0a97f420d63" +} diff --git a/.sqlx/query-bb5d85fbd4f8cfa89b452e7cb1d359a8568238ea6b2bdc2f48933369609f00ac.json b/.sqlx/query-bb5d85fbd4f8cfa89b452e7cb1d359a8568238ea6b2bdc2f48933369609f00ac.json new file mode 100644 index 0000000000..4eb0915b44 --- /dev/null +++ b/.sqlx/query-bb5d85fbd4f8cfa89b452e7cb1d359a8568238ea6b2bdc2f48933369609f00ac.json @@ -0,0 +1,14 @@ +{ + "db_name": "PostgreSQL", + "query": "DELETE FROM group_client_traffic_policy WHERE group_id = $1", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Int8" + ] + }, + "nullable": [] + }, + "hash": "bb5d85fbd4f8cfa89b452e7cb1d359a8568238ea6b2bdc2f48933369609f00ac" +} diff --git a/.sqlx/query-c355c9c90181e0da78602ec73c4cbcc6a6226737625c37347d4d907703552d78.json b/.sqlx/query-c355c9c90181e0da78602ec73c4cbcc6a6226737625c37347d4d907703552d78.json new file mode 100644 index 0000000000..1bef7bdca8 --- /dev/null +++ b/.sqlx/query-c355c9c90181e0da78602ec73c4cbcc6a6226737625c37347d4d907703552d78.json @@ -0,0 +1,39 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT gctp.group_id, gctp.client_traffic_policy \"client_traffic_policy: ClientTrafficPolicy\" FROM group_client_traffic_policy gctp JOIN group_user gu ON gu.group_id = gctp.group_id WHERE gu.user_id = $1", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "group_id", + "type_info": "Int8" + }, + { + "ordinal": 1, + "name": "client_traffic_policy: ClientTrafficPolicy", + "type_info": { + "Custom": { + "name": "client_traffic_policy", + "kind": { + "Enum": [ + "none", + "disable_all_traffic", + "force_all_traffic" + ] + } + } + } + } + ], + "parameters": { + "Left": [ + "Int8" + ] + }, + "nullable": [ + false, + false + ] + }, + "hash": "c355c9c90181e0da78602ec73c4cbcc6a6226737625c37347d4d907703552d78" +} diff --git a/.sqlx/query-eefacbad5c1a6edcb81c569bc287d864722cff2cf1ea01507edf40ba12cdbd69.json b/.sqlx/query-eefacbad5c1a6edcb81c569bc287d864722cff2cf1ea01507edf40ba12cdbd69.json new file mode 100644 index 0000000000..dcc9537ad3 --- /dev/null +++ b/.sqlx/query-eefacbad5c1a6edcb81c569bc287d864722cff2cf1ea01507edf40ba12cdbd69.json @@ -0,0 +1,22 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT id FROM \"group\" WHERE id = ANY($1)", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Int8" + } + ], + "parameters": { + "Left": [ + "Int8Array" + ] + }, + "nullable": [ + false + ] + }, + "hash": "eefacbad5c1a6edcb81c569bc287d864722cff2cf1ea01507edf40ba12cdbd69" +} diff --git a/crates/defguard_core/src/db/models/activity_log/metadata.rs b/crates/defguard_core/src/db/models/activity_log/metadata.rs index 9a87e2ea9e..c74fe07cb6 100644 --- a/crates/defguard_core/src/db/models/activity_log/metadata.rs +++ b/crates/defguard_core/src/db/models/activity_log/metadata.rs @@ -18,7 +18,7 @@ use crate::{ enterprise::db::models::{ activity_log_stream::{ActivityLogStream, ActivityLogStreamType}, api_tokens::ApiToken, - enterprise_settings::EnterpriseSettings, + enterprise_settings::EnterpriseSettingsInfo, openid_provider::{DirectorySyncTarget, DirectorySyncUserBehavior, OpenIdProvider}, snat::UserSnatBinding, }, @@ -357,8 +357,8 @@ pub struct SettingsUpdateMetadata { #[derive(Serialize)] pub struct EnterpriseSettingsUpdateMetadata { - pub before: EnterpriseSettings, - pub after: EnterpriseSettings, + pub before: EnterpriseSettingsInfo, + pub after: EnterpriseSettingsInfo, } #[derive(Serialize)] diff --git a/crates/defguard_core/src/enterprise/db/models/enterprise_settings.rs b/crates/defguard_core/src/enterprise/db/models/enterprise_settings.rs index e251f12908..f27ef08604 100644 --- a/crates/defguard_core/src/enterprise/db/models/enterprise_settings.rs +++ b/crates/defguard_core/src/enterprise/db/models/enterprise_settings.rs @@ -1,4 +1,4 @@ -use defguard_common::db::models::Settings; +use defguard_common::db::{Id, models::Settings}; use sqlx::{PgExecutor, Type, query, query_as}; use struct_patch::Patch; @@ -19,6 +19,33 @@ pub struct EnterpriseSettings { pub display_password_reset: bool, } +#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)] +pub struct GroupClientTrafficPolicies { + pub none: Vec, + pub disable_all_traffic: Vec, + pub force_all_traffic: Vec, +} + +#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)] +pub struct EnterpriseSettingsInfo { + #[serde(flatten)] + pub settings: EnterpriseSettings, + pub group_client_traffic_policies: GroupClientTrafficPolicies, +} + +impl EnterpriseSettingsInfo { + #[must_use] + pub fn new( + settings: EnterpriseSettings, + group_client_traffic_policies: GroupClientTrafficPolicies, + ) -> Self { + Self { + settings, + group_client_traffic_policies, + } + } +} + // We want to be conscious of what the defaults are here #[allow(clippy::derivable_impls)] impl Default for EnterpriseSettings { @@ -107,3 +134,85 @@ pub enum ClientTrafficPolicy { /// Clients are forced to route all traffic through the VPN. ForceAllTraffic, } + +/// Resolves group policies over the instance-level policy. +/// +/// A configured group policy takes precedence over the instance policy. When a user belongs to +/// multiple groups, disabling all traffic takes precedence over forcing all traffic. An explicit +/// `None` group policy takes precedence over the instance policy when no restrictive group policy +/// is present. +#[must_use] +pub fn resolve_client_traffic_policy( + instance_policy: ClientTrafficPolicy, + group_policies: impl IntoIterator, +) -> ClientTrafficPolicy { + let mut has_force_all_traffic = false; + let mut has_none = false; + + for policy in group_policies { + match policy { + ClientTrafficPolicy::DisableAllTraffic => { + return ClientTrafficPolicy::DisableAllTraffic; + } + ClientTrafficPolicy::ForceAllTraffic => has_force_all_traffic = true, + ClientTrafficPolicy::None => has_none = true, + } + } + + if has_force_all_traffic { + ClientTrafficPolicy::ForceAllTraffic + } else if has_none { + ClientTrafficPolicy::None + } else { + instance_policy + } +} + +#[cfg(test)] +mod tests { + use super::{ClientTrafficPolicy, resolve_client_traffic_policy}; + + #[test] + fn instance_policy_is_used_without_group_overrides() { + assert_eq!( + resolve_client_traffic_policy(ClientTrafficPolicy::ForceAllTraffic, []), + ClientTrafficPolicy::ForceAllTraffic + ); + } + + #[test] + fn group_policy_overrides_instance_policy() { + assert_eq!( + resolve_client_traffic_policy( + ClientTrafficPolicy::DisableAllTraffic, + [ClientTrafficPolicy::ForceAllTraffic] + ), + ClientTrafficPolicy::ForceAllTraffic + ); + } + + #[test] + fn disable_all_traffic_wins_conflicting_group_policies() { + assert_eq!( + resolve_client_traffic_policy( + ClientTrafficPolicy::ForceAllTraffic, + [ + ClientTrafficPolicy::ForceAllTraffic, + ClientTrafficPolicy::DisableAllTraffic, + ] + ), + ClientTrafficPolicy::DisableAllTraffic + ); + } + + #[test] + fn explicit_none_group_policy_overrides_instance_policy() { + assert_eq!( + resolve_client_traffic_policy( + ClientTrafficPolicy::ForceAllTraffic, + [ClientTrafficPolicy::None] + ), + ClientTrafficPolicy::None + ); + } +} diff --git a/crates/defguard_core/src/enterprise/db/models/group_client_traffic_policy.rs b/crates/defguard_core/src/enterprise/db/models/group_client_traffic_policy.rs new file mode 100644 index 0000000000..d8e7bf89cb --- /dev/null +++ b/crates/defguard_core/src/enterprise/db/models/group_client_traffic_policy.rs @@ -0,0 +1,126 @@ +use defguard_common::db::Id; +use sqlx::{FromRow, PgConnection, PgExecutor, query, query_as}; + +use super::enterprise_settings::{ClientTrafficPolicy, GroupClientTrafficPolicies}; + +#[derive(Clone, Debug, FromRow, PartialEq)] +/// A traffic policy assigned to a single group. +pub struct GroupClientTrafficPolicy { + pub group_id: Id, + pub client_traffic_policy: ClientTrafficPolicy, +} + +impl GroupClientTrafficPolicy { + pub async fn all<'e, E>(executor: E) -> sqlx::Result> + where + E: PgExecutor<'e>, + { + query_as!( + Self, + "SELECT group_id, \ + client_traffic_policy \"client_traffic_policy: ClientTrafficPolicy\" \ + FROM group_client_traffic_policy ORDER BY group_id" + ) + .fetch_all(executor) + .await + } + + /// Converts database assignments into the API's policy-grouped representation. + #[must_use] + pub fn grouped(policies: Vec) -> GroupClientTrafficPolicies { + let mut grouped = GroupClientTrafficPolicies::default(); + for policy in policies { + match policy.client_traffic_policy { + ClientTrafficPolicy::None => grouped.none.push(policy.group_id), + ClientTrafficPolicy::DisableAllTraffic => { + grouped.disable_all_traffic.push(policy.group_id); + } + ClientTrafficPolicy::ForceAllTraffic => { + grouped.force_all_traffic.push(policy.group_id); + } + } + } + grouped + } + + /// Replaces all group policy assignments within an existing transaction. + pub async fn replace_all( + transaction: &mut PgConnection, + policies: &GroupClientTrafficPolicies, + ) -> sqlx::Result<()> { + query!("DELETE FROM group_client_traffic_policy") + .execute(&mut *transaction) + .await?; + for (group_ids, policy) in [ + (&policies.none, ClientTrafficPolicy::None), + ( + &policies.disable_all_traffic, + ClientTrafficPolicy::DisableAllTraffic, + ), + ( + &policies.force_all_traffic, + ClientTrafficPolicy::ForceAllTraffic, + ), + ] { + for &group_id in group_ids { + Self::upsert(&mut *transaction, group_id, policy).await?; + } + } + Ok(()) + } + + /// Returns policy assignments for groups belonging to a user. + pub async fn find_by_user_id<'e, E>(executor: E, user_id: Id) -> sqlx::Result> + where + E: PgExecutor<'e>, + { + query_as!( + Self, + "SELECT gctp.group_id, \ + gctp.client_traffic_policy \"client_traffic_policy: ClientTrafficPolicy\" \ + FROM group_client_traffic_policy gctp \ + JOIN group_user gu ON gu.group_id = gctp.group_id \ + WHERE gu.user_id = $1", + user_id + ) + .fetch_all(executor) + .await + } + + pub async fn upsert<'e, E>( + executor: E, + group_id: Id, + policy: ClientTrafficPolicy, + ) -> sqlx::Result + where + E: PgExecutor<'e>, + { + query_as!( + Self, + "INSERT INTO group_client_traffic_policy \ + (group_id, client_traffic_policy) \ + VALUES ($1, $2) \ + ON CONFLICT (group_id) DO UPDATE SET \ + client_traffic_policy = EXCLUDED.client_traffic_policy \ + RETURNING group_id, \ + client_traffic_policy \"client_traffic_policy: ClientTrafficPolicy\"", + group_id, + policy as ClientTrafficPolicy + ) + .fetch_one(executor) + .await + } + + pub async fn delete<'e, E>(executor: E, group_id: Id) -> sqlx::Result<()> + where + E: PgExecutor<'e>, + { + query!( + "DELETE FROM group_client_traffic_policy WHERE group_id = $1", + group_id + ) + .execute(executor) + .await?; + Ok(()) + } +} diff --git a/crates/defguard_core/src/enterprise/db/models/mod.rs b/crates/defguard_core/src/enterprise/db/models/mod.rs index e835e1f65a..e327617405 100644 --- a/crates/defguard_core/src/enterprise/db/models/mod.rs +++ b/crates/defguard_core/src/enterprise/db/models/mod.rs @@ -3,5 +3,6 @@ pub mod activity_log_stream; pub mod api_tokens; pub mod device_posture; pub mod enterprise_settings; +pub mod group_client_traffic_policy; pub mod openid_provider; pub mod snat; diff --git a/crates/defguard_core/src/enterprise/handlers/enterprise_settings.rs b/crates/defguard_core/src/enterprise/handlers/enterprise_settings.rs index 0f374dbe96..dce7f36e55 100644 --- a/crates/defguard_core/src/enterprise/handlers/enterprise_settings.rs +++ b/crates/defguard_core/src/enterprise/handlers/enterprise_settings.rs @@ -1,16 +1,95 @@ +use std::collections::HashSet; + use axum::{Json, extract::State, http::StatusCode}; -use defguard_common::types::proxy::ProxyControlMessage; +use defguard_common::{db::Id, types::proxy::ProxyControlMessage}; +use sqlx::{PgConnection, PgPool, query_scalar}; use struct_patch::Patch; use super::LicenseInfo; use crate::{ appstate::AppState, auth::{AdminRole, SessionInfo}, - enterprise::db::models::enterprise_settings::{EnterpriseSettings, EnterpriseSettingsPatch}, + enterprise::db::models::{ + enterprise_settings::{ + ClientTrafficPolicy, EnterpriseSettings, EnterpriseSettingsInfo, + EnterpriseSettingsPatch, GroupClientTrafficPolicies, + }, + group_client_traffic_policy::GroupClientTrafficPolicy, + }, + error::WebError, events::{ApiEvent, ApiEventType, ApiRequestContext}, handlers::{ApiResponse, ApiResult}, }; +#[derive(Deserialize)] +/// Request payload for partially updating enterprise settings and group policies. +pub struct EnterpriseSettingsPatchRequest { + #[serde(flatten)] + pub settings: EnterpriseSettingsPatch, + pub group_client_traffic_policies: Option, +} + +fn policy_assignments(policies: &GroupClientTrafficPolicies) -> Vec<(Id, ClientTrafficPolicy)> { + policies + .none + .iter() + .map(|&id| (id, ClientTrafficPolicy::None)) + .chain( + policies + .disable_all_traffic + .iter() + .map(|&id| (id, ClientTrafficPolicy::DisableAllTraffic)), + ) + .chain( + policies + .force_all_traffic + .iter() + .map(|&id| (id, ClientTrafficPolicy::ForceAllTraffic)), + ) + .collect() +} + +async fn validate_policy_assignments( + transaction: &mut PgConnection, + policies: &GroupClientTrafficPolicies, +) -> Result<(), WebError> { + let assignments = policy_assignments(policies); + let mut group_ids = HashSet::with_capacity(assignments.len()); + for (group_id, _) in &assignments { + if !group_ids.insert(*group_id) { + return Err(WebError::BadRequest( + "A group cannot be assigned to multiple client traffic policies.".into(), + )); + } + } + + let group_ids = group_ids.into_iter().collect::>(); + if group_ids.is_empty() { + return Ok(()); + } + let existing_ids: HashSet = + query_scalar!("SELECT id FROM \"group\" WHERE id = ANY($1)", &group_ids) + .fetch_all(&mut *transaction) + .await? + .into_iter() + .collect(); + if existing_ids.len() != group_ids.len() { + return Err(WebError::BadRequest( + "One or more client traffic policy groups do not exist.".into(), + )); + } + Ok(()) +} + +async fn settings_info( + pool: &PgPool, + settings: EnterpriseSettings, +) -> Result { + let group_policies = + GroupClientTrafficPolicy::grouped(GroupClientTrafficPolicy::all(pool).await?); + Ok(EnterpriseSettingsInfo::new(settings, group_policies)) +} + pub async fn get_enterprise_settings( session: SessionInfo, State(appstate): State, @@ -24,7 +103,10 @@ pub async fn get_enterprise_settings( "User {} retrieved enterprise settings", session.user.username ); - Ok(ApiResponse::json(settings, StatusCode::OK)) + Ok(ApiResponse::json( + settings_info(&appstate.pool, settings).await?, + StatusCode::OK, + )) } pub async fn patch_enterprise_settings( @@ -32,13 +114,16 @@ pub async fn patch_enterprise_settings( _admin: AdminRole, State(appstate): State, session: SessionInfo, - Json(data): Json, + Json(data): Json, ) -> ApiResult { debug!( "Admin {} patching enterprise settings.", session.user.username, ); - let mut settings = EnterpriseSettings::get(&appstate.pool).await?; + let mut transaction = appstate.pool.begin().await?; + let mut settings = EnterpriseSettings::get(&mut *transaction).await?; + let old_group_policies = + GroupClientTrafficPolicy::grouped(GroupClientTrafficPolicy::all(&mut *transaction).await?); // snapshot for audit event let old_settings = settings.clone(); @@ -46,9 +131,23 @@ pub async fn patch_enterprise_settings( let old_display_password_reset = old_settings.display_password_reset; let old_display_download_step = old_settings.display_download_step; - settings.apply(data); - settings.save(&appstate.pool).await?; - info!("Admin {} patched settings.", session.user.username); + settings.apply(data.settings); + let group_policies = if let Some(group_policies) = data.group_client_traffic_policies { + validate_policy_assignments(&mut transaction, &group_policies).await?; + GroupClientTrafficPolicy::replace_all(&mut transaction, &group_policies).await?; + group_policies + } else { + old_group_policies.clone() + }; + settings.save(&mut *transaction).await?; + transaction.commit().await?; + + let before = EnterpriseSettingsInfo::new(old_settings, old_group_policies); + let after = EnterpriseSettingsInfo::new(settings.clone(), group_policies); + info!( + "Admin {} patched enterprise settings.", + session.user.username + ); appstate.emit_event(ApiEvent { context: ApiRequestContext::new( @@ -57,10 +156,7 @@ pub async fn patch_enterprise_settings( None::, "web".into(), ), - event: Box::new(ApiEventType::EnterpriseSettingsUpdated { - before: old_settings, - after: settings.clone(), - }), + event: Box::new(ApiEventType::EnterpriseSettingsUpdated { before, after }), })?; // Broadcast updated public settings to proxies only if they changed. diff --git a/crates/defguard_core/src/events.rs b/crates/defguard_core/src/events.rs index 4fb8a2855b..5edbe8a916 100644 --- a/crates/defguard_core/src/events.rs +++ b/crates/defguard_core/src/events.rs @@ -17,7 +17,7 @@ use crate::{ activity_log_stream::ActivityLogStream, api_tokens::ApiToken, device_posture::{DevicePosture, DevicePostureSnapshot}, - enterprise_settings::EnterpriseSettings, + enterprise_settings::EnterpriseSettingsInfo, openid_provider::OpenIdProvider, snat::UserSnatBinding, }, @@ -251,8 +251,8 @@ pub enum ApiEventType { }, SettingsDefaultBrandingRestored, EnterpriseSettingsUpdated { - before: EnterpriseSettings, - after: EnterpriseSettings, + before: EnterpriseSettingsInfo, + after: EnterpriseSettingsInfo, }, GroupsBulkAssigned { users: Vec>, diff --git a/crates/defguard_core/src/grpc/mod.rs b/crates/defguard_core/src/grpc/mod.rs index 283897785f..73394af13a 100644 --- a/crates/defguard_core/src/grpc/mod.rs +++ b/crates/defguard_core/src/grpc/mod.rs @@ -10,7 +10,7 @@ use defguard_common::{ config::server_config, db::{ Id, - models::{Settings, WireguardNetwork, wireguard::ServiceLocationMode}, + models::{Settings, User, WireguardNetwork, wireguard::ServiceLocationMode}, }, types::UrlParseError, }; @@ -24,7 +24,10 @@ use crate::{ enterprise::{ LicenseFeature, db::models::{ - enterprise_settings::{ClientTrafficPolicy, EnterpriseSettings}, + enterprise_settings::{ + ClientTrafficPolicy, EnterpriseSettings, resolve_client_traffic_policy, + }, + group_client_traffic_policy::GroupClientTrafficPolicy, openid_provider::OpenIdProvider, }, has_enterprise_access, is_business_license_active, @@ -163,13 +166,33 @@ pub struct InstanceInfo { openid_display_name: Option, } +#[derive(Debug, thiserror::Error)] +/// Errors that can occur while building client instance information. +pub enum InstanceInfoBuildError { + #[error("failed to load enterprise settings: {0}")] + Database(#[from] sqlx::Error), + #[error("failed to parse instance URL: {0}")] + UrlParse(#[from] UrlParseError), +} + impl InstanceInfo { - pub fn new>( - settings: Settings, - username: S, - enterprise_settings: &EnterpriseSettings, + /// Builds client instance information with the effective user traffic policy. + pub async fn build( + pool: &PgPool, + settings: &Settings, + user: &User, openid_provider: Option>, - ) -> Result { + ) -> Result { + let enterprise_settings = EnterpriseSettings::get(pool).await?; + let client_traffic_policy = if is_business_license_active() { + let group_policies = GroupClientTrafficPolicy::find_by_user_id(pool, user.id) + .await? + .into_iter() + .map(|policy| policy.client_traffic_policy); + resolve_client_traffic_policy(enterprise_settings.client_traffic_policy, group_policies) + } else { + enterprise_settings.client_traffic_policy + }; let openid_display_name = openid_provider .as_ref() .map(|provider| provider.display_name.clone()) @@ -178,11 +201,11 @@ impl InstanceInfo { let proxy_url = settings.proxy_public_url()?; Ok(Self { id: settings.uuid, - name: settings.instance_name, + name: settings.instance_name.clone(), url, proxy_url, - username: username.into(), - client_traffic_policy: enterprise_settings.client_traffic_policy, + username: user.username.clone(), + client_traffic_policy, enterprise_enabled: is_business_license_active(), openid_display_name, }) diff --git a/crates/defguard_core/src/grpc/utils.rs b/crates/defguard_core/src/grpc/utils.rs index 7f1fd478e0..5b8d040180 100644 --- a/crates/defguard_core/src/grpc/utils.rs +++ b/crates/defguard_core/src/grpc/utils.rs @@ -24,9 +24,7 @@ use tonic::Status; use super::InstanceInfo; use crate::{ device_access::build_device_config, - enterprise::db::models::{ - enterprise_settings::EnterpriseSettings, openid_provider::OpenIdProvider, - }, + enterprise::db::models::openid_provider::OpenIdProvider, grpc::{client_version::ClientFeature, should_prevent_service_location_usage}, }; @@ -48,11 +46,6 @@ pub async fn build_device_config_response( Status::internal(format!("unexpected error: {err}")) })?; - let enterprise_settings = EnterpriseSettings::get(pool).await.map_err(|err| { - error!("Failed to get enterprise settings: {err}"); - Status::internal(format!("unexpected error: {err}")) - })?; - let mut configs = Vec::new(); let user = User::find_by_id(pool, device.user_id) .await @@ -231,16 +224,12 @@ pub async fn build_device_config_response( user.username, user.id, device.name, device.id ); - let instance_info = InstanceInfo::new( - settings, - &user.username, - &enterprise_settings, - openid_provider, - ) - .map_err(|err| { - error!("Failed to build instance info: {err}"); - Status::internal(format!("unexpected error: {err}")) - })?; + let instance_info = InstanceInfo::build(pool, &settings, &user, openid_provider) + .await + .map_err(|err| { + error!("Failed to build instance info: {err}"); + Status::internal(format!("unexpected error: {err}")) + })?; Ok(DeviceConfigResponse { device: Some(device.into()), diff --git a/crates/defguard_core/tests/integration/api/enterprise_settings.rs b/crates/defguard_core/tests/integration/api/enterprise_settings.rs index 4704557464..03fbc21851 100644 --- a/crates/defguard_core/tests/integration/api/enterprise_settings.rs +++ b/crates/defguard_core/tests/integration/api/enterprise_settings.rs @@ -1,9 +1,11 @@ use std::time::Duration; -use defguard_common::types::proxy::ProxyControlMessage; +use defguard_common::{db::models::group::Group, types::proxy::ProxyControlMessage}; use defguard_core::{ enterprise::{ - db::models::enterprise_settings::{ClientTrafficPolicy, EnterpriseSettings}, + db::models::enterprise_settings::{ + ClientTrafficPolicy, EnterpriseSettings, EnterpriseSettingsInfo, + }, license::{get_cached_license, set_cached_license}, }, events::ApiEventType, @@ -523,13 +525,13 @@ async fn test_display_flags_round_trip(_: PgPoolOptions, options: PgConnectOptio // Read back and verify the values persisted let response = client.get("/api/v1/settings_enterprise").send().await; assert_eq!(response.status(), StatusCode::OK); - let body: EnterpriseSettings = response.json().await; + let body: EnterpriseSettingsInfo = response.json().await; assert!( - !body.display_download_step, + !body.settings.display_download_step, "display_download_step should be false" ); assert!( - !body.display_password_reset, + !body.settings.display_password_reset, "display_password_reset should be false" ); @@ -551,13 +553,13 @@ async fn test_display_flags_round_trip(_: PgPoolOptions, options: PgConnectOptio // Read back and verify let response = client.get("/api/v1/settings_enterprise").send().await; assert_eq!(response.status(), StatusCode::OK); - let body: EnterpriseSettings = response.json().await; + let body: EnterpriseSettingsInfo = response.json().await; assert!( - body.display_download_step, + body.settings.display_download_step, "display_download_step should be true" ); assert!( - body.display_password_reset, + body.settings.display_password_reset, "display_password_reset should be true" ); } @@ -734,3 +736,142 @@ async fn test_public_settings_broadcast_on_save(_: PgPoolOptions, options: PgCon } } } + +#[sqlx::test] +async fn test_group_client_traffic_policies_are_saved_and_validated( + _: PgPoolOptions, + options: PgConnectOptions, +) { + let pool = setup_pool(options).await; + let (client, _) = make_test_client(pool.clone()).await; + let auth = Auth::new("admin", "pass123"); + assert_eq!( + client + .post("/api/v1/auth") + .json(&auth) + .send() + .await + .status(), + StatusCode::OK + ); + exceed_enterprise_limits(&client).await; + + let allow_choice = Group::new("allow-choice").save(&pool).await.unwrap(); + let disable = Group::new("disable").save(&pool).await.unwrap(); + + let response = client + .patch("/api/v1/settings_enterprise") + .json(&json!({ + "client_traffic_policy": "force_all_traffic", + "group_client_traffic_policies": { + "none": [allow_choice.id], + "disable_all_traffic": [disable.id], + "force_all_traffic": [] + } + })) + .send() + .await; + assert_eq!(response.status(), StatusCode::OK); + + let response = client.get("/api/v1/settings_enterprise").send().await; + assert_eq!(response.status(), StatusCode::OK); + let settings: EnterpriseSettingsInfo = response.json().await; + assert_eq!( + settings.group_client_traffic_policies.none, + vec![allow_choice.id] + ); + assert_eq!( + settings.group_client_traffic_policies.disable_all_traffic, + vec![disable.id] + ); + assert!( + settings + .group_client_traffic_policies + .force_all_traffic + .is_empty() + ); + + let license = get_cached_license().clone(); + set_cached_license(None); + let response = client.get("/api/v1/settings_enterprise").send().await; + assert_eq!(response.status(), StatusCode::OK); + let settings: EnterpriseSettingsInfo = response.json().await; + assert_eq!( + settings.group_client_traffic_policies.none, + vec![allow_choice.id] + ); + assert_eq!( + settings.group_client_traffic_policies.disable_all_traffic, + vec![disable.id] + ); + set_cached_license(license); + + let response = client + .patch("/api/v1/settings_enterprise") + .json(&json!({"display_download_step": false})) + .send() + .await; + assert_eq!(response.status(), StatusCode::OK); + + let response = client.get("/api/v1/settings_enterprise").send().await; + assert_eq!(response.status(), StatusCode::OK); + let settings: EnterpriseSettingsInfo = response.json().await; + assert_eq!( + settings.group_client_traffic_policies.none, + vec![allow_choice.id] + ); + assert_eq!( + settings.group_client_traffic_policies.disable_all_traffic, + vec![disable.id] + ); + + let response = client + .patch("/api/v1/settings_enterprise") + .json(&json!({ + "group_client_traffic_policies": { + "none": [allow_choice.id], + "disable_all_traffic": [allow_choice.id], + "force_all_traffic": [] + } + })) + .send() + .await; + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + + let response = client.get("/api/v1/settings_enterprise").send().await; + assert_eq!(response.status(), StatusCode::OK); + let settings: EnterpriseSettingsInfo = response.json().await; + assert_eq!( + settings.group_client_traffic_policies.none, + vec![allow_choice.id] + ); + assert_eq!( + settings.group_client_traffic_policies.disable_all_traffic, + vec![disable.id] + ); + + let response = client + .patch("/api/v1/settings_enterprise") + .json(&json!({ + "group_client_traffic_policies": { + "none": [999999], + "disable_all_traffic": [], + "force_all_traffic": [] + } + })) + .send() + .await; + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + + let response = client.get("/api/v1/settings_enterprise").send().await; + assert_eq!(response.status(), StatusCode::OK); + let settings: EnterpriseSettingsInfo = response.json().await; + assert_eq!( + settings.group_client_traffic_policies.none, + vec![allow_choice.id] + ); + assert_eq!( + settings.group_client_traffic_policies.disable_all_traffic, + vec![disable.id] + ); +} diff --git a/crates/defguard_event_logger/src/tests/mod.rs b/crates/defguard_event_logger/src/tests/mod.rs index 9807d8babd..8fa5caec82 100644 --- a/crates/defguard_event_logger/src/tests/mod.rs +++ b/crates/defguard_event_logger/src/tests/mod.rs @@ -23,7 +23,7 @@ use defguard_core::{ activity_log_stream::{ActivityLogStream, ActivityLogStreamType}, api_tokens::ApiToken, device_posture::{DevicePosture, DevicePostureSnapshot}, - enterprise_settings::EnterpriseSettings, + enterprise_settings::EnterpriseSettingsInfo, openid_provider::{ DirectorySyncTarget, DirectorySyncUserBehavior, OpenIdProvider, OpenIdProviderKind, }, @@ -852,8 +852,8 @@ fn api_event_cases() -> Vec { EventTestCase { name: "EnterpriseSettingsUpdated", message: api_message(ApiEventType::EnterpriseSettingsUpdated { - before: EnterpriseSettings::default(), - after: EnterpriseSettings::default(), + before: EnterpriseSettingsInfo::default(), + after: EnterpriseSettingsInfo::default(), }), event_type: EventType::EnterpriseSettingsUpdated, module: ActivityLogModule::Defguard, diff --git a/crates/defguard_proxy_manager/src/servers/enrollment.rs b/crates/defguard_proxy_manager/src/servers/enrollment.rs index 3851a08cb1..ead524de5f 100644 --- a/crates/defguard_proxy_manager/src/servers/enrollment.rs +++ b/crates/defguard_proxy_manager/src/servers/enrollment.rs @@ -217,16 +217,12 @@ impl EnrollmentServer { Status::internal(format!("unexpected error: {err}")) })?; let smtp_configured = settings.smtp_configured(); - let instance_info = InstanceInfo::new( - settings, - &user.username, - &enterprise_settings, - openid_provider, - ) - .map_err(|err| { - error!("Failed to create instance info: {err}"); - Status::internal("unexpected error") - })?; + let instance_info = InstanceInfo::build(&self.pool, &settings, &user, openid_provider) + .await + .map_err(|err| { + error!("Failed to create instance info: {err}"); + Status::internal("unexpected error") + })?; debug!("Instance info {instance_info:?}"); debug!( @@ -956,16 +952,12 @@ impl EnrollmentServer { Status::internal(format!("unexpected error: {err}")) })?; - let instance_info = InstanceInfo::new( - settings, - &user.username, - &enterprise_settings, - openid_provider, - ) - .map_err(|err| { - error!("Failed to create instance info: {err}"); - Status::internal("unexpected error") - })?; + let instance_info = InstanceInfo::build(&self.pool, &settings, &user, openid_provider) + .await + .map_err(|err| { + error!("Failed to create instance info: {err}"); + Status::internal("unexpected error") + })?; let response = DeviceConfigResponse { device: Some(device.clone().into()), diff --git a/crates/defguard_proxy_manager/src/tests/proxy_manager/handler/polling.rs b/crates/defguard_proxy_manager/src/tests/proxy_manager/handler/polling.rs index ca9bc9cd4b..8de15dd370 100644 --- a/crates/defguard_proxy_manager/src/tests/proxy_manager/handler/polling.rs +++ b/crates/defguard_proxy_manager/src/tests/proxy_manager/handler/polling.rs @@ -1,9 +1,20 @@ -use defguard_core::device_access::join_device_to_all_networks; +use defguard_common::db::models::group::Group; +use defguard_core::{ + device_access::join_device_to_all_networks, + enterprise::db::models::{ + enterprise_settings::ClientTrafficPolicy, + group_client_traffic_policy::GroupClientTrafficPolicy, + }, + grpc::utils::build_device_config_response, +}; use defguard_proto::{ - client_types::InstanceInfoRequest, + client_types, proxy::{CoreRequest, core_request, core_response}, }; -use sqlx::postgres::{PgConnectOptions, PgPoolOptions}; +use sqlx::{ + PgPool, + postgres::{PgConnectOptions, PgPoolOptions}, +}; use super::support::{ assert_error_response, clear_test_license, complete_proxy_handshake, create_device_for_user, @@ -12,6 +23,121 @@ use super::support::{ }; use crate::tests::common::HandlerTestContext; +async fn poll_client_traffic_policy(context: &mut HandlerTestContext, token: &str) -> i32 { + context.mock_proxy().send_request(CoreRequest { + id: 20, + device_info: None, + payload: Some(core_request::Payload::InstanceInfo( + client_types::InstanceInfoRequest { + token: token.to_owned(), + }, + )), + }); + + let response = context.mock_proxy_mut().recv_outbound().await; + match response.payload { + Some(core_response::Payload::InstanceInfo(info)) => info + .device_config + .and_then(|config| config.instance) + .and_then(|instance| instance.client_traffic_policy) + .expect("InstanceInfo should contain a client traffic policy"), + other => panic!( + "expected InstanceInfo response, got: {:?}", + other.as_ref().map(std::mem::discriminant) + ), + } +} + +async fn set_global_client_traffic_policy(pool: &PgPool, policy: &str) { + sqlx::query( + "UPDATE \"enterprisesettings\" SET client_traffic_policy = $1::client_traffic_policy WHERE id = 1", + ) + .bind(policy) + .execute(pool) + .await + .expect("failed to update global client traffic policy"); +} + +#[sqlx::test] +async fn test_client_traffic_policy_is_resolved_for_instance_info( + _: PgPoolOptions, + options: PgConnectOptions, +) { + let mut context = HandlerTestContext::new(options).await; + complete_proxy_handshake(&mut context).await; + set_test_license_business(); + + let _network = create_network(&context.pool).await; + let (user, device) = create_user_with_device(&context.pool).await; + let token = create_polling_token(&context.pool, device.id).await; + let force_group = Group::new("force-policy") + .save(&context.pool) + .await + .unwrap(); + let disable_group = Group::new("disable-policy") + .save(&context.pool) + .await + .unwrap(); + + set_global_client_traffic_policy(&context.pool, "disable_all_traffic").await; + assert_eq!( + poll_client_traffic_policy(&mut context, &token).await, + client_types::ClientTrafficPolicy::DisableAllTraffic as i32 + ); + + user.add_to_group(&context.pool, &force_group) + .await + .unwrap(); + GroupClientTrafficPolicy::upsert( + &context.pool, + force_group.id, + ClientTrafficPolicy::ForceAllTraffic, + ) + .await + .unwrap(); + assert_eq!( + poll_client_traffic_policy(&mut context, &token).await, + client_types::ClientTrafficPolicy::ForceAllTraffic as i32 + ); + + GroupClientTrafficPolicy::upsert(&context.pool, force_group.id, ClientTrafficPolicy::None) + .await + .unwrap(); + assert_eq!( + poll_client_traffic_policy(&mut context, &token).await, + client_types::ClientTrafficPolicy::None as i32 + ); + + user.add_to_group(&context.pool, &disable_group) + .await + .unwrap(); + GroupClientTrafficPolicy::upsert( + &context.pool, + disable_group.id, + ClientTrafficPolicy::DisableAllTraffic, + ) + .await + .unwrap(); + assert_eq!( + poll_client_traffic_policy(&mut context, &token).await, + client_types::ClientTrafficPolicy::DisableAllTraffic as i32 + ); + + clear_test_license(); + let response = build_device_config_response(&context.pool, device, None, None) + .await + .expect("failed to build device config without a license"); + assert_eq!( + response + .instance + .expect("device config should contain instance info") + .client_traffic_policy, + Some(client_types::ClientTrafficPolicy::None as i32) + ); + + context.finish().await.expect_server_finished().await; +} + #[sqlx::test] async fn test_polling_returns_updated_device_config(_: PgPoolOptions, options: PgConnectOptions) { let mut context = HandlerTestContext::new(options).await; @@ -27,9 +153,11 @@ async fn test_polling_returns_updated_device_config(_: PgPoolOptions, options: P context.mock_proxy().send_request(CoreRequest { id: 10, device_info: None, - payload: Some(core_request::Payload::InstanceInfo(InstanceInfoRequest { - token: token_str.clone(), - })), + payload: Some(core_request::Payload::InstanceInfo( + client_types::InstanceInfoRequest { + token: token_str.clone(), + }, + )), }); let response = context.mock_proxy_mut().recv_outbound().await; @@ -64,9 +192,9 @@ async fn test_polling_requires_business_license(_: PgPoolOptions, options: PgCon context.mock_proxy().send_request(CoreRequest { id: 11, device_info: None, - payload: Some(core_request::Payload::InstanceInfo(InstanceInfoRequest { - token: token_str, - })), + payload: Some(core_request::Payload::InstanceInfo( + client_types::InstanceInfoRequest { token: token_str }, + )), }); let response = context.mock_proxy_mut().recv_outbound().await; @@ -90,9 +218,11 @@ async fn test_polling_invalid_token_returns_error(_: PgPoolOptions, options: PgC context.mock_proxy().send_request(CoreRequest { id: 12, device_info: None, - payload: Some(core_request::Payload::InstanceInfo(InstanceInfoRequest { - token: "this-token-does-not-exist-00000000".to_owned(), - })), + payload: Some(core_request::Payload::InstanceInfo( + client_types::InstanceInfoRequest { + token: "this-token-does-not-exist-00000000".to_owned(), + }, + )), }); let response = context.mock_proxy_mut().recv_outbound().await; @@ -127,9 +257,9 @@ async fn test_polling_inactive_user_returns_error(_: PgPoolOptions, options: PgC context.mock_proxy().send_request(CoreRequest { id: 13, device_info: None, - payload: Some(core_request::Payload::InstanceInfo(InstanceInfoRequest { - token: token_str, - })), + payload: Some(core_request::Payload::InstanceInfo( + client_types::InstanceInfoRequest { token: token_str }, + )), }); let response = context.mock_proxy_mut().recv_outbound().await; @@ -159,9 +289,11 @@ async fn test_polling_reflects_network_changes(_: PgPoolOptions, options: PgConn context.mock_proxy().send_request(CoreRequest { id: 14, device_info: None, - payload: Some(core_request::Payload::InstanceInfo(InstanceInfoRequest { - token: token_str.clone(), - })), + payload: Some(core_request::Payload::InstanceInfo( + client_types::InstanceInfoRequest { + token: token_str.clone(), + }, + )), }); let first_response = context.mock_proxy_mut().recv_outbound().await; let first_info = match &first_response.payload { @@ -182,9 +314,11 @@ async fn test_polling_reflects_network_changes(_: PgPoolOptions, options: PgConn context.mock_proxy().send_request(CoreRequest { id: 15, device_info: None, - payload: Some(core_request::Payload::InstanceInfo(InstanceInfoRequest { - token: token_str.clone(), - })), + payload: Some(core_request::Payload::InstanceInfo( + client_types::InstanceInfoRequest { + token: token_str.clone(), + }, + )), }); let second_response = context.mock_proxy_mut().recv_outbound().await; let second_info = match &second_response.payload { diff --git a/migrations/20260723000000_[2.1.0]_group_client_traffic_policy.down.sql b/migrations/20260723000000_[2.1.0]_group_client_traffic_policy.down.sql new file mode 100644 index 0000000000..f5219e6850 --- /dev/null +++ b/migrations/20260723000000_[2.1.0]_group_client_traffic_policy.down.sql @@ -0,0 +1 @@ +DROP TABLE group_client_traffic_policy; diff --git a/migrations/20260723000000_[2.1.0]_group_client_traffic_policy.up.sql b/migrations/20260723000000_[2.1.0]_group_client_traffic_policy.up.sql new file mode 100644 index 0000000000..9a648e6b70 --- /dev/null +++ b/migrations/20260723000000_[2.1.0]_group_client_traffic_policy.up.sql @@ -0,0 +1,4 @@ +CREATE TABLE group_client_traffic_policy ( + group_id bigint PRIMARY KEY REFERENCES "group"(id) ON DELETE CASCADE, + client_traffic_policy client_traffic_policy NOT NULL +); diff --git a/web/messages/en/groups.json b/web/messages/en/groups.json index b89bf6757d..4c638fa95e 100644 --- a/web/messages/en/groups.json +++ b/web/messages/en/groups.json @@ -7,6 +7,10 @@ "groups_col_name": "Group name", "groups_col_users_count": "Added users", "groups_col_type": "Type", + "groups_col_traffic_policy": "Traffic policy", + "groups_traffic_policy_none": "No limitations", + "groups_traffic_policy_disable_all": "Disable all traffic", + "groups_traffic_policy_force_all": "Force all traffic", "groups_col_locations": "Used in locations", "groups_type_admin": "Admin", "groups_type_user": "User" diff --git a/web/messages/en/settings.json b/web/messages/en/settings.json index 4e1cf227b5..36bfe04b2d 100644 --- a/web/messages/en/settings.json +++ b/web/messages/en/settings.json @@ -320,12 +320,20 @@ "settings_client_section_traffic_policy_title": "Client traffic policy", "settings_client_traffic_policy_description_title": "Client traffic rules", "settings_client_traffic_policy_description": "Specify the conditions that determine how traffic should behave in the application.", - "settings_client_traffic_policy_none_title": "None", + "settings_client_traffic_policy_none_title": "No limitation", "settings_client_traffic_policy_none_content": "When this option is enabled, users will be able to select all routing options.", "settings_client_traffic_policy_disable_all_title": "Disable all traffic", "settings_client_traffic_policy_disable_all_content": "When this option is enabled, users will not be able to route all traffic through the VPN.", "settings_client_traffic_policy_force_all_title": "Force all traffic", "settings_client_traffic_policy_force_all_content": "When this option is enabled, the users will always route all traffic through the VPN.", + "settings_client_traffic_policy_group_title": "Group-based policies", + "settings_client_traffic_policy_group_description": "Define the groups that should use their own traffic rules instead of the global traffic policy. Any groups not included below will use the global traffic policy.", + "settings_client_traffic_policy_group_none_content": "When this option is enabled, users in groups will be able to select all routing options.", + "settings_client_traffic_policy_group_disable_all_title": "Disable all traffic", + "settings_client_traffic_policy_group_disable_all_content": "When this option is enabled, users in groups will not be able to route all traffic through the VPN.", + "settings_client_traffic_policy_group_force_all_title": "Force all traffic", + "settings_client_traffic_policy_group_force_all_content": "When this option is enabled, the users in groups will always route all traffic through the VPN.", + "settings_client_traffic_policy_edit_groups": "Edit groups", "settings_gateway_notifications_title": "Gateway notifications", "settings_gateway_notifications_subtitle": "Here you can manage email notifications.", "settings_notifications_gateway_card_content": "Configure admin email notifications for gateway disconnect and reconnect events, and set the inactivity threshold that triggers disconnect notifications.", diff --git a/web/src/pages/GroupsPage/GroupsPage.tsx b/web/src/pages/GroupsPage/GroupsPage.tsx index fc63ac999f..8f7f7a942d 100644 --- a/web/src/pages/GroupsPage/GroupsPage.tsx +++ b/web/src/pages/GroupsPage/GroupsPage.tsx @@ -5,6 +5,7 @@ import { SizedBox } from '../../shared/defguard-ui/components/SizedBox/SizedBox' import { ThemeSpacing } from '../../shared/defguard-ui/types'; import { TablePageLayout } from '../../shared/layout/TablePageLayout/TablePageLayout'; import { + getEnterpriseSettingsQueryOptions, getGroupsInfoQueryOptions, getLocationsQueryOptions, getUsersOverviewQueryOptions, @@ -14,14 +15,27 @@ import { CEGroupModal } from './modals/CEGroupModal/CEGroupModal'; export const GroupsPage = () => { const { data: groups } = useSuspenseQuery(getGroupsInfoQueryOptions); + const { data: enterpriseSettings } = useSuspenseQuery( + getEnterpriseSettingsQueryOptions, + ); const { data: locations } = useSuspenseQuery(getLocationsQueryOptions); const { data: users } = useSuspenseQuery(getUsersOverviewQueryOptions); + const groupClientTrafficPolicies = enterpriseSettings.group_client_traffic_policies ?? { + none: [], + disable_all_traffic: [], + force_all_traffic: [], + }; return ( <> - + diff --git a/web/src/pages/GroupsPage/components/GroupsTable/GroupsTable.tsx b/web/src/pages/GroupsPage/components/GroupsTable/GroupsTable.tsx index d250102443..e0c9859bbd 100644 --- a/web/src/pages/GroupsPage/components/GroupsTable/GroupsTable.tsx +++ b/web/src/pages/GroupsPage/components/GroupsTable/GroupsTable.tsx @@ -7,7 +7,12 @@ import { import { useMemo, useState } from 'react'; import { m } from '../../../../paraglide/messages'; import api from '../../../../shared/api/api'; -import type { GroupInfo, NetworkLocation, User } from '../../../../shared/api/types'; +import type { + GroupClientTrafficPolicies, + GroupInfo, + NetworkLocation, + User, +} from '../../../../shared/api/types'; import { Badge } from '../../../../shared/defguard-ui/components/Badge/Badge'; import { Button } from '../../../../shared/defguard-ui/components/Button/Button'; import type { MenuItemProps } from '../../../../shared/defguard-ui/components/Menu/types'; @@ -23,6 +28,7 @@ import { ModalName } from '../../../../shared/hooks/modalControls/modalTypes'; type Props = { groups: GroupInfo[]; + groupClientTrafficPolicies: GroupClientTrafficPolicies; locations: NetworkLocation[]; users: User[]; }; @@ -31,7 +37,12 @@ type RowData = GroupInfo; const columnHelper = createColumnHelper(); -export const GroupsTable = ({ groups, locations, users }: Props) => { +export const GroupsTable = ({ + groups, + groupClientTrafficPolicies, + locations, + users, +}: Props) => { const [search, setSearch] = useState(''); const reservedNames = useMemo(() => groups.map((g) => g.name), [groups]); @@ -82,6 +93,27 @@ export const GroupsTable = ({ groups, locations, users }: Props) => { ), }), + columnHelper.display({ + id: 'traffic_policy', + minSize: 200, + header: m.groups_col_traffic_policy(), + cell: (info) => { + const groupId = info.row.original.id; + let policy = '-'; + if (groupClientTrafficPolicies.none.includes(groupId)) { + policy = m.groups_traffic_policy_none(); + } else if (groupClientTrafficPolicies.disable_all_traffic.includes(groupId)) { + policy = m.groups_traffic_policy_disable_all(); + } else if (groupClientTrafficPolicies.force_all_traffic.includes(groupId)) { + policy = m.groups_traffic_policy_force_all(); + } + return ( + + {policy} + + ); + }, + }), columnHelper.accessor('vpn_locations', { minSize: 350, header: m.groups_col_locations(), @@ -157,7 +189,7 @@ export const GroupsTable = ({ groups, locations, users }: Props) => { }, }), ], - [locations, reservedNames, users], + [groupClientTrafficPolicies, locations, reservedNames, users], ); const table = useReactTable({ diff --git a/web/src/pages/settings/SettingsClientPage/SettingsClientPage.tsx b/web/src/pages/settings/SettingsClientPage/SettingsClientPage.tsx index 1af2e693db..989415f24d 100644 --- a/web/src/pages/settings/SettingsClientPage/SettingsClientPage.tsx +++ b/web/src/pages/settings/SettingsClientPage/SettingsClientPage.tsx @@ -1,6 +1,9 @@ import { useMutation, useQuery, useSuspenseQuery } from '@tanstack/react-query'; import { Link } from '@tanstack/react-router'; -import { ClientTrafficPolicy } from '../../../shared/api/types'; +import { + ClientTrafficPolicy, + type GroupClientTrafficPolicies, +} from '../../../shared/api/types'; import { Breadcrumbs } from '../../../shared/components/Breadcrumbs/Breadcrumbs'; import { ContextualHelpKey, @@ -8,15 +11,19 @@ import { } from '../../../shared/components/ContextualHelp'; import { DescriptionBlock } from '../../../shared/components/DescriptionBlock/DescriptionBlock'; import { Page } from '../../../shared/components/Page/Page'; +import type { SelectionOption } from '../../../shared/components/SelectionSection/type'; +import { SelectMultiple } from '../../../shared/components/SelectMultiple/SelectMultiple'; import { SettingsCard } from '../../../shared/components/SettingsCard/SettingsCard'; import { SettingsHeader } from '../../../shared/components/SettingsHeader/SettingsHeader'; import { SettingsLayout } from '../../../shared/components/SettingsLayout/SettingsLayout'; import { Divider } from '../../../shared/defguard-ui/components/Divider/Divider'; +import { Icon, type IconKindValue } from '../../../shared/defguard-ui/components/Icon'; import { MarkedSection } from '../../../shared/defguard-ui/components/MarkedSection/MarkedSection'; -import { ThemeSpacing } from '../../../shared/defguard-ui/types'; +import { ThemeSpacing, ThemeVariable } from '../../../shared/defguard-ui/types'; import { isPresent } from '../../../shared/defguard-ui/utils/isPresent'; import { getEnterpriseSettingsQueryOptions, + getGroupsInfoQueryOptions, getLicenseInfoQueryOptions, } from '../../../shared/query'; import './style.scss'; @@ -31,9 +38,7 @@ import { Button } from '../../../shared/defguard-ui/components/Button/Button'; import { Snackbar } from '../../../shared/defguard-ui/providers/snackbar/snackbar'; import { useAppForm } from '../../../shared/form'; import { formChangeLogic } from '../../../shared/formLogic'; -import { openModal } from '../../../shared/hooks/modalControls/modalsSubjects'; -import { ModalName } from '../../../shared/hooks/modalControls/modalTypes'; -import { canUseBusinessFeature } from '../../../shared/utils/license'; +import { canUseBusinessFeature, licenseActionCheck } from '../../../shared/utils/license'; const breadcrumbs = [ @@ -45,7 +50,7 @@ const breadcrumbs = [ ]; export const SettingsClientPage = () => { - const { data: license, isFetched } = useQuery(getLicenseInfoQueryOptions); + const { data: license } = useQuery(getLicenseInfoQueryOptions); return ( @@ -56,7 +61,11 @@ export const SettingsClientPage = () => { icon="user" title={m.settings_client_title()} subtitle={m.settings_client_subtitle()} - badgeProps={!isPresent(license) && isFetched ? businessBadgeProps : undefined} + badgeProps={ + license !== undefined && !canUseBusinessFeature(license).result + ? businessBadgeProps + : undefined + } /> }> @@ -70,15 +79,102 @@ const formSchema = z.object({ admin_device_management: z.boolean(), only_client_activation: z.boolean(), client_traffic_policy: z.enum(ClientTrafficPolicy), + group_client_traffic_policies: z.object({ + none: z.array(z.number()), + disable_all_traffic: z.array(z.number()), + force_all_traffic: z.array(z.number()), + }), }); type FormFields = z.infer; +type GroupPolicy = keyof GroupClientTrafficPolicies; + +const emptyGroupClientTrafficPolicies = { + none: [], + disable_all_traffic: [], + force_all_traffic: [], +}; + +type GroupPolicyRowProps = { + canEdit: boolean; + content: string; + title: string; + icon: IconKindValue; + options: SelectionOption[]; + selected: number[]; + onSelectionChange: (value: number[]) => void; + onEditUnavailable: () => void; +}; + +const getSelectedGroupsCounterText = (count: number) => { + if (count === 1) return m.location_access_selected_group_count_one({ count }); + return m.location_access_selected_group_count_other({ count }); +}; + +const getAvailableGroupOptions = ( + options: SelectionOption[], + policy: GroupPolicy, + policies: GroupClientTrafficPolicies, +) => { + const assignedToOtherPolicy = new Set( + Object.entries(policies) + .filter(([key]) => key !== policy) + .flatMap(([, groupIds]) => groupIds), + ); + + return options.filter((option) => !assignedToOtherPolicy.has(option.id)); +}; + +const GroupPolicyRow = ({ + content, + title, + icon, + canEdit, + onSelectionChange, + options, + selected, + onEditUnavailable, +}: GroupPolicyRowProps) => ( +
+ +
+

{title}

+

{content}

+ {canEdit ? ( + {}} + options={options} + selected={new Set(selected)} + toggleValue={false} + /> + ) : ( + + )} +
+
+); + const Content = () => { const { data: licenseInfo } = useSuspenseQuery(getLicenseInfoQueryOptions); const { data: settings } = useSuspenseQuery(getEnterpriseSettingsQueryOptions); + const { data: groups } = useSuspenseQuery(getGroupsInfoQueryOptions); const noLicense = !isPresent(licenseInfo); + const canUseTrafficPolicies = canUseBusinessFeature(licenseInfo).result; + const groupClientTrafficPolicies = canUseTrafficPolicies + ? (settings.group_client_traffic_policies ?? emptyGroupClientTrafficPolicies) + : emptyGroupClientTrafficPolicies; const { mutateAsync: patchSettings } = useMutation({ mutationFn: api.settings.patchEnterpriseSettings, @@ -97,14 +193,24 @@ const Content = () => { return { admin_device_management: settings.admin_device_management, only_client_activation: settings.only_client_activation, - client_traffic_policy: settings.client_traffic_policy, + client_traffic_policy: canUseTrafficPolicies + ? settings.client_traffic_policy + : ClientTrafficPolicy.None, + group_client_traffic_policies: groupClientTrafficPolicies, }; }, [ settings.admin_device_management, settings.client_traffic_policy, settings.only_client_activation, + canUseTrafficPolicies, + groupClientTrafficPolicies, ]); + const groupOptions = groups.map>((group) => ({ + id: group.id, + label: group.name, + })); + const form = useAppForm({ defaultValues, validationLogic: formChangeLogic, @@ -113,17 +219,13 @@ const Content = () => { onChange: formSchema, }, onSubmit: async ({ value }) => { - if (!licenseInfo) return; - // only expire error is possible here - const { result } = canUseBusinessFeature(licenseInfo); - if (result) { - await patchSettings(value); - form.reset(value); - } else { - openModal(ModalName.LicenseExpired, { - licenseTier: licenseInfo?.tier, - }); + const licenseCheck = canUseBusinessFeature(licenseInfo); + if (!licenseCheck.result) { + licenseActionCheck(licenseCheck, () => {}); + return; } + await patchSettings(value); + form.reset(value); }, }); @@ -174,7 +276,7 @@ const Content = () => { {(field) => ( { {(field) => ( { {(field) => ( { )} + + +

{m.settings_client_traffic_policy_group_title()}

+

+ {m.settings_client_traffic_policy_group_description()} +

+ state.values.group_client_traffic_policies} + > + {(policies) => ( + <> + + form.setFieldValue('group_client_traffic_policies', { + ...policies, + none, + }) + } + options={getAvailableGroupOptions(groupOptions, 'none', policies)} + selected={policies.none} + onEditUnavailable={() => + licenseActionCheck(canUseBusinessFeature(licenseInfo), () => {}) + } + /> + + form.setFieldValue('group_client_traffic_policies', { + ...policies, + disable_all_traffic, + }) + } + options={getAvailableGroupOptions( + groupOptions, + 'disable_all_traffic', + policies, + )} + selected={policies.disable_all_traffic} + onEditUnavailable={() => + licenseActionCheck(canUseBusinessFeature(licenseInfo), () => {}) + } + /> + + form.setFieldValue('group_client_traffic_policies', { + ...policies, + force_all_traffic, + }) + } + options={getAvailableGroupOptions( + groupOptions, + 'force_all_traffic', + policies, + )} + selected={policies.force_all_traffic} + onEditUnavailable={() => + licenseActionCheck(canUseBusinessFeature(licenseInfo), () => {}) + } + /> + + )} + +
({ isDefault: s.isDefaultValue || s.isPristine, diff --git a/web/src/pages/settings/SettingsClientPage/style.scss b/web/src/pages/settings/SettingsClientPage/style.scss index df76f2d176..3135ec7058 100644 --- a/web/src/pages/settings/SettingsClientPage/style.scss +++ b/web/src/pages/settings/SettingsClientPage/style.scss @@ -7,5 +7,45 @@ h3 { font: var(--t-body-primary-600); } + + .group-policy-description { + color: var(--fg-muted); + font: var(--t-body-sm-400); + } + + .group-policy-row { + display: flex; + align-items: flex-start; + column-gap: var(--spacing-md); + } + + .group-policy-row-content { + display: flex; + flex: 1; + flex-flow: column; + row-gap: var(--spacing-sm); + } + + .group-policy-title { + font: var(--t-body-primary-500); + } + + .group-policy-content { + color: var(--fg-neutral); + font: var(--t-body-sm-400); + } + + .select-multiple-edit { + align-self: flex-start; + display: inline-flex; + align-items: center; + gap: var(--spacing-sm); + padding: 0; + border: 0; + background-color: transparent; + color: var(--fg-action); + font: var(--t-body-sm-500); + cursor: pointer; + } } } diff --git a/web/src/pages/settings/SettingsIndexPage/tabs/SettingsGeneralTab.tsx b/web/src/pages/settings/SettingsIndexPage/tabs/SettingsGeneralTab.tsx index 2aef6c2f1c..81b0005bba 100644 --- a/web/src/pages/settings/SettingsIndexPage/tabs/SettingsGeneralTab.tsx +++ b/web/src/pages/settings/SettingsIndexPage/tabs/SettingsGeneralTab.tsx @@ -11,6 +11,7 @@ import { SectionSelect } from '../../../../shared/defguard-ui/components/Section import { SizedBox } from '../../../../shared/defguard-ui/components/SizedBox/SizedBox'; import { ThemeSpacing } from '../../../../shared/defguard-ui/types'; import { getLicenseInfoQueryOptions } from '../../../../shared/query'; +import { canUseBusinessFeature } from '../../../../shared/utils/license'; export const SettingsGeneralTab = () => { const navigate = useNavigate(); @@ -34,7 +35,11 @@ export const SettingsGeneralTab = () => { image="behavior" title={m.settings_breadcrumb_client_behavior()} content={m.settings_general_section_client_behavior_content()} - badgeProps={licenseInfo === null ? businessBadgeProps : undefined} + badgeProps={ + licenseInfo !== undefined && !canUseBusinessFeature(licenseInfo).result + ? businessBadgeProps + : undefined + } onClick={() => { navigate({ to: '/settings/client' }); }} diff --git a/web/src/shared/api/types.ts b/web/src/shared/api/types.ts index f110e19964..7f903fea8b 100644 --- a/web/src/shared/api/types.ts +++ b/web/src/shared/api/types.ts @@ -951,12 +951,19 @@ export const ClientTrafficPolicy = { export type ClientTrafficPolicyValue = (typeof ClientTrafficPolicy)[keyof typeof ClientTrafficPolicy]; +export interface GroupClientTrafficPolicies { + none: number[]; + disable_all_traffic: number[]; + force_all_traffic: number[]; +} + export interface SettingsEnterprise { admin_device_management: boolean; client_traffic_policy: ClientTrafficPolicyValue; only_client_activation: boolean; display_download_step: boolean; display_password_reset: boolean; + group_client_traffic_policies: GroupClientTrafficPolicies; } export type ApiDevicePostureOsRule = diff --git a/web/src/shared/defguard-ui b/web/src/shared/defguard-ui index 95e7fd8f66..693c5ecf22 160000 --- a/web/src/shared/defguard-ui +++ b/web/src/shared/defguard-ui @@ -1 +1 @@ -Subproject commit 95e7fd8f66c31f0cac247359a63fae3177b5dbf1 +Subproject commit 693c5ecf2252d33ae9d45bbf36574e6108565944