diff --git a/new-ui/src/shared/components/LocationCard/components/LocationCardConnectionTiles/LocationCardConnectionTiles.tsx b/new-ui/src/shared/components/LocationCard/components/LocationCardConnectionTiles/LocationCardConnectionTiles.tsx index fa31c17de..a0331475b 100644 --- a/new-ui/src/shared/components/LocationCard/components/LocationCardConnectionTiles/LocationCardConnectionTiles.tsx +++ b/new-ui/src/shared/components/LocationCard/components/LocationCardConnectionTiles/LocationCardConnectionTiles.tsx @@ -2,7 +2,12 @@ import './style.scss'; import clsx from 'clsx'; import { useMemo } from 'react'; import { useAppData } from '../../../../providers/AppDataContext'; -import { type LocationInfo, LocationMfaMode } from '../../../../rust-api/types'; +import { + ClientTrafficPolicy, + type InstanceInfo, + type LocationInfo, + LocationMfaMode, +} from '../../../../rust-api/types'; import { isPresent } from '../../../../utils/isPresent'; import { mfaToText } from '../../../../utils/mfa'; import { BoxIcon } from '../../../BoxIcon/BoxIcon'; @@ -11,11 +16,16 @@ import { Icon, IconKind } from '../../../Icon'; interface Props { variant: 'compact' | 'full'; location: LocationInfo; + instance?: InstanceInfo; } -export const LocationCardConnectionTiles = ({ location, variant }: Props) => { +export const LocationCardConnectionTiles = ({ location, instance, variant }: Props) => { const { connectionMfaMethod } = useAppData(); + const routeAllTraffic = + location.route_all_traffic || + instance?.client_traffic_policy === ClientTrafficPolicy.ForceAllTraffic; + const mfaMethod = useMemo(() => { const key = `${location.connection_type.toLowerCase()}-${location.id}`; const method = connectionMfaMethod[key]; @@ -30,7 +40,7 @@ export const LocationCardConnectionTiles = ({ location, variant }: Props) => {

Allowed traffic

- {location.route_all_traffic ? 'All traffic' : 'Predefined traffic'} + {routeAllTraffic ? 'All traffic' : 'Predefined traffic'}

{location.location_mfa_mode !== LocationMfaMode.Disabled && diff --git a/new-ui/src/shared/components/LocationCard/views/ConnectedView/ConnectedView.tsx b/new-ui/src/shared/components/LocationCard/views/ConnectedView/ConnectedView.tsx index bfade887e..7b1b0c30e 100644 --- a/new-ui/src/shared/components/LocationCard/views/ConnectedView/ConnectedView.tsx +++ b/new-ui/src/shared/components/LocationCard/views/ConnectedView/ConnectedView.tsx @@ -7,12 +7,16 @@ import { LocationCardConnectionTiles } from '../../components/LocationCardConnec import { useLocationCardContext } from '../../context/context'; export const ConnectedView = () => { - const { location } = useLocationCardContext(); + const { location, instance } = useLocationCardContext(); return (
- + diff --git a/new-ui/src/shared/components/OverviewLocationCard/OverviewLocationCard.tsx b/new-ui/src/shared/components/OverviewLocationCard/OverviewLocationCard.tsx index a76ee16c2..a224906f8 100644 --- a/new-ui/src/shared/components/OverviewLocationCard/OverviewLocationCard.tsx +++ b/new-ui/src/shared/components/OverviewLocationCard/OverviewLocationCard.tsx @@ -128,7 +128,11 @@ export const OverviewLocationCard = ({ location, instance }: Props) => {
{location.active && ( - + )} {!location.active && ( diff --git a/src-tauri/client-cli/src/commands/list.rs b/src-tauri/client-cli/src/commands/list.rs index 3d3eb2f32..9228ebf9f 100644 --- a/src-tauri/client-cli/src/commands/list.rs +++ b/src-tauri/client-cli/src/commands/list.rs @@ -1,6 +1,11 @@ use std::collections::HashMap; -use defguard_core::database::models::{instance::Instance, location::Location, tunnel::Tunnel, Id}; +use defguard_core::database::models::{ + instance::{ClientTrafficPolicy, Instance}, + location::Location, + tunnel::Tunnel, + Id, +}; use serde_json::{json, Value}; use crate::{ @@ -44,10 +49,10 @@ impl CommandOutput for ListResult { } fn json(&self) -> Value { - let instance_names = self + let instances_by_id = self .instances .iter() - .map(|i| (i.id, i.name.clone())) + .map(|i| (i.id, i)) .collect::>(); let instances = self @@ -63,15 +68,25 @@ impl CommandOutput for ListResult { let locations = self .locations .iter() - .map(|l| LocationEntry { - id: l.id, - name: l.name.clone(), - instance: instance_names.get(&l.instance_id).cloned(), - address: l.address.clone(), - endpoint: l.endpoint.clone(), - mfa_enabled: Some(l.mfa_enabled()), - mfa_method: Some(mfa_label(l.mfa_method).to_string()), - route_all_traffic: Some(l.route_all_traffic), + .map(|l| { + let instance = instances_by_id.get(&l.instance_id); + let route_all_traffic = match instance + .map_or(&ClientTrafficPolicy::None, |i| &i.client_traffic_policy) + { + ClientTrafficPolicy::None => l.route_all_traffic, + ClientTrafficPolicy::DisableAllTraffic => false, + ClientTrafficPolicy::ForceAllTraffic => true, + }; + LocationEntry { + id: l.id, + name: l.name.clone(), + instance: instance.map(|i| i.name.clone()), + address: l.address.clone(), + endpoint: l.endpoint.clone(), + mfa_enabled: Some(l.mfa_enabled()), + mfa_method: Some(mfa_label(l.mfa_method).to_string()), + route_all_traffic: Some(route_all_traffic), + } }) .collect::>(); @@ -131,11 +146,18 @@ fn format_list_table( )); for location in locations { let mfa = if location.mfa_enabled() { "yes" } else { "no" }; - let route_label = if location.route_all_traffic { + let route_all_traffic = match instance.client_traffic_policy { + ClientTrafficPolicy::None => location.route_all_traffic, + ClientTrafficPolicy::DisableAllTraffic => false, + ClientTrafficPolicy::ForceAllTraffic => true, + }; + + let route_label = if route_all_traffic { "All-traffic" } else { "Predefined" }; + lines.push(format!( " {:>4} {:3} {route_label:<11}", location.id, location.name, location.address, location.endpoint diff --git a/src-tauri/client-cli/src/commands/location.rs b/src-tauri/client-cli/src/commands/location.rs index 503e07a7c..139ac4dc7 100644 --- a/src-tauri/client-cli/src/commands/location.rs +++ b/src-tauri/client-cli/src/commands/location.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use defguard_core::database::models::{ - instance::Instance, + instance::{ClientTrafficPolicy, Instance}, location::{Location, LocationMfaMethod}, Id, }; @@ -20,15 +20,23 @@ const MIN_INST_COL_WIDTH: usize = 8; pub(crate) async fn handle_list(state: &State) -> Result { let locations = Location::all(&state.pool, false).await?; - let instance_names = Instance::all(&state.pool) + let instance_details = Instance::all(&state.pool) .await? .into_iter() - .map(|inst| (inst.id, inst.name)) + .map(|instance| { + ( + instance.id, + InstanceDetails { + name: instance.name, + client_traffic_policy: instance.client_traffic_policy, + }, + ) + }) .collect::>(); Ok(LocationListResult { locations, - instance_names, + instance_details, }) } @@ -93,6 +101,11 @@ pub async fn handle_show( let ResolvedTarget::Location(location) = &target else { return Err(CliError::NotFound(format!("Location '{name}' not found"))); }; + let client_traffic_policy = Instance::find_by_id(&state.pool, location.instance_id) + .await? + .map_or(ClientTrafficPolicy::None, |instance| { + instance.client_traffic_policy + }); Ok(LocationShowResult { name: location.name.clone(), @@ -102,7 +115,11 @@ pub async fn handle_show( allowed_ips: location.allowed_ips.clone(), dns: location.dns.clone(), mfa_method: mfa_label(location.mfa_method).to_string(), - route_all_traffic: location.route_all_traffic, + route_all_traffic: match client_traffic_policy { + ClientTrafficPolicy::None => location.route_all_traffic, + ClientTrafficPolicy::DisableAllTraffic => false, + ClientTrafficPolicy::ForceAllTraffic => true, + }, keepalive_interval: location.keepalive_interval, }) } @@ -127,9 +144,14 @@ pub(crate) fn mfa_label(method: Option) -> &'static str { } } +pub(crate) struct InstanceDetails { + pub name: String, + pub client_traffic_policy: ClientTrafficPolicy, +} + pub struct LocationListResult { pub locations: Vec>, - pub instance_names: HashMap, + pub instance_details: HashMap, } impl CommandOutput for LocationListResult { @@ -137,7 +159,7 @@ impl CommandOutput for LocationListResult { if self.locations.is_empty() { "No locations configured. Use the desktop app to enroll an instance first.".to_string() } else { - format_location_list_table(&self.locations, &self.instance_names) + format_location_list_table(&self.locations, &self.instance_details) } } @@ -145,15 +167,26 @@ impl CommandOutput for LocationListResult { let locations = self .locations .iter() - .map(|l| LocationEntry { - id: l.id, - name: l.name.clone(), - instance: self.instance_names.get(&l.instance_id).cloned(), - address: l.address.clone(), - endpoint: l.endpoint.clone(), - mfa_enabled: None, - mfa_method: Some(mfa_label(l.mfa_method).to_string()), - route_all_traffic: Some(l.route_all_traffic), + .map(|l| { + let details = self.instance_details.get(&l.instance_id); + let route_all_traffic = match details + .map_or(&ClientTrafficPolicy::None, |details| { + &details.client_traffic_policy + }) { + ClientTrafficPolicy::None => l.route_all_traffic, + ClientTrafficPolicy::DisableAllTraffic => false, + ClientTrafficPolicy::ForceAllTraffic => true, + }; + LocationEntry { + id: l.id, + name: l.name.clone(), + instance: details.map(|details| details.name.clone()), + address: l.address.clone(), + endpoint: l.endpoint.clone(), + mfa_enabled: None, + mfa_method: Some(mfa_label(l.mfa_method).to_string()), + route_all_traffic: Some(route_all_traffic), + } }) .collect::>(); json!({ "locations": locations }) @@ -162,7 +195,7 @@ impl CommandOutput for LocationListResult { fn format_location_list_table( locations: &[Location], - instance_names: &HashMap, + instance_details: &HashMap, ) -> String { let name_col_width = locations .iter() @@ -179,9 +212,9 @@ fn format_location_list_table( let inst_col_width = locations .iter() .filter_map(|l| { - instance_names + instance_details .get(&l.instance_id) - .map(std::string::String::len) + .map(|details| details.name.len()) }) .max() .unwrap_or(MIN_INST_COL_WIDTH) @@ -192,22 +225,34 @@ fn format_location_list_table( "ID", "LOCATION", "ADDRESS", "ENDPOINT", "INSTANCE", "MFA", "Routing" )]; for location in locations { - let instance = instance_names - .get(&location.instance_id) - .map_or("?", String::as_str); + let details = instance_details.get(&location.instance_id); + + let instance_name = details.map_or("?", |instance| instance.name.as_str()); + let instance_traffic_policy = details.map_or(&ClientTrafficPolicy::None, |instance| { + &instance.client_traffic_policy + }); + + let route_all_traffic = match instance_traffic_policy { + ClientTrafficPolicy::None => location.route_all_traffic, + ClientTrafficPolicy::DisableAllTraffic => false, + ClientTrafficPolicy::ForceAllTraffic => true, + }; + + let route_label = if route_all_traffic { + "All-traffic" + } else { + "Predefined" + }; + lines.push(format!( " {:>4} {:3} {:>11}", location.id, location.name, location.address, location.endpoint, - instance, + instance_name, mfa_label(location.mfa_method), - if location.route_all_traffic { - "All-traffic" - } else { - "Predefined" - } + route_label )); } lines.join("\n") @@ -322,11 +367,21 @@ mod tests { } } + fn make_instance_details( + name: &str, + client_traffic_policy: ClientTrafficPolicy, + ) -> InstanceDetails { + InstanceDetails { + name: name.to_string(), + client_traffic_policy, + } + } + #[test] fn test_list_human_empty() { let result = LocationListResult { locations: Vec::new(), - instance_names: HashMap::new(), + instance_details: HashMap::new(), }; assert_eq!( result.human(), @@ -337,11 +392,11 @@ mod tests { #[test] fn test_list_human_with_data() { let loc = make_location(1, 10, "office", "1.2.3.4:51820", false); - let mut names = HashMap::new(); - names.insert(10, "acme".to_string()); + let mut instance_details = HashMap::new(); + instance_details.insert(10, make_instance_details("acme", ClientTrafficPolicy::None)); let result = LocationListResult { locations: vec![loc], - instance_names: names, + instance_details, }; let s = result.human(); assert!(s.contains("ID")); @@ -350,11 +405,46 @@ mod tests { assert!(s.contains("1.2.3.4:51820")); } + fn routing_column( + route_all_traffic: bool, + client_traffic_policy: ClientTrafficPolicy, + ) -> String { + let mut location = make_location(1, 10, "office", "1.2.3.4:51820", false); + location.route_all_traffic = route_all_traffic; + let mut instance_details = HashMap::new(); + instance_details.insert(10, make_instance_details("acme", client_traffic_policy)); + LocationListResult { + locations: vec![location], + instance_details, + } + .human() + } + + #[test] + fn test_list_human_force_all_traffic_overrides_location() { + let table = routing_column(false, ClientTrafficPolicy::ForceAllTraffic); + assert!(table.contains("All-traffic")); + assert!(!table.contains("Predefined")); + } + + #[test] + fn test_list_human_disable_all_traffic_overrides_location() { + let table = routing_column(true, ClientTrafficPolicy::DisableAllTraffic); + assert!(table.contains("Predefined")); + assert!(!table.contains("All-traffic")); + } + + #[test] + fn test_list_human_no_policy_keeps_location_setting() { + assert!(routing_column(true, ClientTrafficPolicy::None).contains("All-traffic")); + assert!(routing_column(false, ClientTrafficPolicy::None).contains("Predefined")); + } + #[test] fn test_list_json_empty() { let result = LocationListResult { locations: Vec::new(), - instance_names: HashMap::new(), + instance_details: HashMap::new(), }; let json = result.json(); assert_eq!(json["locations"].as_array().unwrap().len(), 0); @@ -364,11 +454,11 @@ mod tests { #[test] fn test_list_json_with_data() { let loc = make_location(1, 10, "office", "1.2.3.4:51820", false); - let mut names = HashMap::new(); - names.insert(10, "acme".to_string()); + let mut instance_details = HashMap::new(); + instance_details.insert(10, make_instance_details("acme", ClientTrafficPolicy::None)); let result = LocationListResult { locations: vec![loc], - instance_names: names, + instance_details, }; let json = result.json(); let locations = json["locations"].as_array().unwrap(); @@ -458,7 +548,7 @@ mod tests { assert_eq!( LocationListResult { locations: Vec::new(), - instance_names: HashMap::new(), + instance_details: HashMap::new(), } .exit_code(), 0