Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion dstack/vmm/src/app.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion dstack/vmm/src/config.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
114 changes: 76 additions & 38 deletions dstack/vmm/src/main_service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<Manifest> {
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<Vec<PortMapping>> {
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::<Result<Vec<_>>>()?;
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<Manifest> {
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,
Expand Down Expand Up @@ -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::<Result<Vec<_>>>()?;
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;
}
Expand Down Expand Up @@ -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()?;
Expand Down
Loading