From b90450412309c366abc155b6f1706ffded1c25bc Mon Sep 17 00:00:00 2001 From: Kevin Wang Date: Wed, 23 Sep 2026 21:51:57 -0700 Subject: [PATCH] fix(vmm): enforce port mapping policy on update_vm --- dstack/vmm/src/app.rs | 2 +- dstack/vmm/src/config.rs | 2 +- dstack/vmm/src/main_service.rs | 114 ++++++++++++++++++++++----------- 3 files changed, 78 insertions(+), 40 deletions(-) diff --git a/dstack/vmm/src/app.rs b/dstack/vmm/src/app.rs index c63fcff74..df3e99281 100644 --- a/dstack/vmm/src/app.rs +++ b/dstack/vmm/src/app.rs @@ -87,7 +87,7 @@ fn signal_pidfd(pid: u32, signal: libc::c_int) -> std::io::Result<()> { } } -#[derive(Deserialize, Serialize, Debug, Clone)] +#[derive(Deserialize, Serialize, Debug, Clone, PartialEq, Eq)] pub struct PortMapping { pub address: IpAddr, pub protocol: Protocol, diff --git a/dstack/vmm/src/config.rs b/dstack/vmm/src/config.rs index 23684737c..0d3d40a3e 100644 --- a/dstack/vmm/src/config.rs +++ b/dstack/vmm/src/config.rs @@ -81,7 +81,7 @@ pub fn load_config_figment(config_file: Option<&str>) -> Figment { load_config("vmm", DEFAULT_CONFIG, config_file, false) } -#[derive(Debug, Clone, Deserialize, Serialize)] +#[derive(Debug, Clone, PartialEq, Eq, Deserialize, Serialize)] #[serde(rename_all = "lowercase")] pub enum Protocol { Tcp, diff --git a/dstack/vmm/src/main_service.rs b/dstack/vmm/src/main_service.rs index f81ff674d..b2c34fe43 100644 --- a/dstack/vmm/src/main_service.rs +++ b/dstack/vmm/src/main_service.rs @@ -29,7 +29,9 @@ use crate::app::{ validate_resolved_networks, App, AttachMode, GpuConfig, GpuSpec, Manifest, PortMapping, VmWorkDir, }; -use crate::config::{CvmConfig, DiskPrealloc, Networking, NetworkingMode, NicNetworking}; +use crate::config::{ + CvmConfig, DiskPrealloc, Networking, NetworkingMode, NicNetworking, PortMappingConfig, +}; fn hex_sha256(data: &str) -> String { use sha2::Digest; @@ -263,42 +265,57 @@ fn validate_unique_port_mappings(mappings: &[PortMapping]) -> Result<()> { Ok(()) } -// Shared function to create manifest from VM configuration -pub fn create_manifest_from_vm_config( - request: VmConfiguration, - cvm_config: &crate::config::CvmConfig, -) -> Result { - validate_label(&request.name)?; - - let pm_cfg = &cvm_config.port_mapping; - if !(request.ports.is_empty() || pm_cfg.enabled) { - bail!("Port mapping is disabled"); - } - let port_map = request - .ports +/// Converts requested port mappings, enforcing `cvm.port_mapping` on every +/// mapping not already in `held`. The web UI resends a VM's full port list on +/// every update, so an existing mapping must stay accepted after the node +/// narrows its policy. +fn port_map_from_proto( + ports: &[rpc::PortMapping], + pm_cfg: &PortMappingConfig, + held: &[PortMapping], +) -> Result> { + let port_map = ports .iter() .map(|p| { - let from = p.host_port.try_into().context("Invalid host port")?; - let to = p.vm_port.try_into().context("Invalid vm port")?; - if !pm_cfg.is_allowed(&p.protocol, from) { - bail!("Port mapping is not allowed for {}:{}", p.protocol, from); - } - let protocol = p.protocol.parse().context("Invalid protocol")?; let address = if !p.host_address.is_empty() { p.host_address.parse().context("Invalid host address")? } else { pm_cfg.address }; - Ok(PortMapping { + let mapping = PortMapping { address, - protocol, - from, - to, + protocol: p.protocol.parse().context("Invalid protocol")?, + from: p.host_port.try_into().context("Invalid host port")?, + to: p.vm_port.try_into().context("Invalid vm port")?, nic_index: p.nic_index.map(|index| index as usize), - }) + }; + if !held.contains(&mapping) { + if !pm_cfg.enabled { + bail!("Port mapping is disabled"); + } + if !pm_cfg.is_allowed(mapping.protocol.as_str(), mapping.from) { + bail!( + "Port mapping is not allowed for {}:{}", + mapping.protocol.as_str(), + mapping.from + ); + } + } + Ok(mapping) }) .collect::>>()?; validate_unique_port_mappings(&port_map)?; + Ok(port_map) +} + +// Shared function to create manifest from VM configuration +pub fn create_manifest_from_vm_config( + request: VmConfiguration, + cvm_config: &crate::config::CvmConfig, +) -> Result { + validate_label(&request.name)?; + + let port_map = port_map_from_proto(&request.ports, &cvm_config.port_mapping, &[])?; let networks = networks_from_vm_config(&request, cvm_config)?; validate_port_mapping_nics( &port_map, @@ -1093,19 +1110,11 @@ impl VmmRpc for RpcHandler { manifest.no_tee = no_tee; } if request.update_ports { - let port_map = request - .ports - .iter() - .map(|p| { - Ok(PortMapping { - address: p.host_address.parse().context("Invalid host address")?, - protocol: p.protocol.parse().context("Invalid protocol")?, - from: p.host_port.try_into().context("Invalid host port")?, - to: p.vm_port.try_into().context("Invalid vm port")?, - nic_index: p.nic_index.map(|index| index as usize), - }) - }) - .collect::>>()?; + let port_map = port_map_from_proto( + &request.ports, + &self.app.config.cvm.port_mapping, + &manifest.port_map, + )?; self.validate_port_mapping_conflicts(Some(&request.id), &port_map)?; manifest.port_map = port_map; } @@ -2857,6 +2866,35 @@ mod tests { } } + #[test] + fn port_map_enforces_node_policy_on_new_mappings_only() { + let port = |host_port: u32| rpc::PortMapping { + protocol: "tcp".into(), + host_port, + vm_port: host_port, + host_address: String::new(), + nic_index: None, + }; + let mut pm_cfg = test_cvm_config().port_mapping; + + pm_cfg.enabled = false; + let err = port_map_from_proto(&[port(8080)], &pm_cfg, &[]).unwrap_err(); + assert!(err.to_string().contains("disabled"), "{err}"); + + pm_cfg.enabled = true; + let held = port_map_from_proto(&[port(8080)], &pm_cfg, &[]).unwrap(); + assert_eq!(held[0].address, pm_cfg.address); + let err = port_map_from_proto(&[port(30000)], &pm_cfg, &[]).unwrap_err(); + assert!(err.to_string().contains("not allowed"), "{err}"); + + pm_cfg.enabled = false; + assert_eq!( + port_map_from_proto(&[port(8080)], &pm_cfg, &held).unwrap(), + held + ); + assert!(port_map_from_proto(&[port(8080), port(8081)], &pm_cfg, &held).is_err()); + } + #[test] fn resolve_volumes_attaches_duplicate_root_once() -> Result<()> { let tmp = tempfile::tempdir()?;