use std::{fmt::Display, net::Ipv4Addr, path::PathBuf, str::FromStr, time::Duration};
use clap::{Args, Parser, Subcommand, ValueEnum};
use serde::{Deserialize, Serialize};
use tracing::Level;
use crate::utils::{VmexecDirs, escape_path, get_live_cid_and_pids_for_vmid};
#[derive(Debug, Clone, PartialEq, ValueEnum)]
pub enum OsType {
Archlinux,
}
#[derive(Debug, Clone, ValueEnum, Serialize, Deserialize)]
pub enum Interactive {
Always,
Never,
Auto,
}
#[derive(Debug, Clone, ValueEnum, Serialize, Deserialize)]
pub enum Tty {
Always,
Never,
Auto,
}
#[derive(Debug, Clone, PartialEq)]
pub enum OsTypeOrImagePath {
OsType(OsType),
ImagePath(PathBuf),
}
impl FromStr for OsTypeOrImagePath {
type Err = String;
fn from_str(src: &str) -> Result<Self, Self::Err> {
if let Ok(os_type) = OsType::from_str(src, true) {
return Ok(Self::OsType(os_type));
} else {
let path = PathBuf::from(src);
if path.is_file() {
return Ok(Self::ImagePath(path));
}
}
let mut err = format!("Could not parse '{src}' as OS type or as an existing file path\n");
let os_types = format!("{:?}", OsType::value_variants());
err.push_str(&format!("Valid OS types are: {}", os_types.to_lowercase()));
Err(err)
}
}
#[derive(Debug, Clone, Args)]
#[group(required = true, multiple = false)]
pub struct ImageSource {
pub os: Option<OsType>,
#[arg(value_parser = parse_existing_pathbuf)]
pub image: Option<PathBuf>,
}
fn parse_existing_pathbuf(src: &str) -> Result<PathBuf, String> {
let path = PathBuf::from(src)
.canonicalize()
.map_err(|e| format!("Failed to canonicalize path: {e}"))?;
Ok(path)
}
fn parse_seconds_to_duration(src: &str) -> Result<Duration, String> {
let sec_int = src
.parse()
.map_err(|_e| format!("Failed to parse '{src}' as an integer"))?;
Ok(Duration::from_secs(sec_int))
}
fn parse_valid_vmid(src: &str) -> Result<String, String> {
let dirs = VmexecDirs::new().unwrap();
if let Ok(cid_pids) = get_live_cid_and_pids_for_vmid(src, &dirs.runs_dir) {
if cid_pids.qemu_pid.is_some() {
return Ok(src.to_string());
}
}
Err("No virtual machine with provided ID found".to_string())
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PublishPort {
pub host_ip: Ipv4Addr,
pub host_port: u32,
pub vm_port: u32,
}
impl Display for PublishPort {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"{}:{}->{}/tcp",
self.host_ip, self.host_port, self.vm_port
)
}
}
impl FromStr for PublishPort {
type Err = String;
fn from_str(src: &str) -> Result<Self, Self::Err> {
let parts: Vec<&str> = src.split(':').collect();
if parts[0].is_empty() {
return Err("Expected format: [[hostip:][hostport]:]vmport".to_string());
}
let (host_ip, host_port, vm_port) = match parts.len() {
1 => {
let host_ip = Ipv4Addr::UNSPECIFIED;
let host_port = parts[0]
.parse()
.map_err(|_| format!("'{}' is not a valid port", parts[0]))?;
let vm_port = parts[0]
.parse()
.map_err(|_| format!("'{}' is not a valid port", parts[0]))?;
(host_ip, host_port, vm_port)
}
2 => {
let host_ip = Ipv4Addr::UNSPECIFIED;
let host_port = parts[0]
.parse()
.map_err(|_| format!("'{}' is not a valid port", parts[0]))?;
let vm_port = parts[1]
.parse()
.map_err(|_| format!("'{}' is not a valid port", parts[1]))?;
(host_ip, host_port, vm_port)
}
3 => {
let host_ip = parts[0]
.parse()
.map_err(|_| format!("'{}' is not a valid IPv4", parts[0]))?;
let vm_port = parts[2]
.parse()
.map_err(|_| format!("'{}' is not a valid port", parts[2]))?;
let host_port = if !parts[1].is_empty() {
parts[1]
.parse()
.map_err(|_| format!("'{}' is not a valid port", parts[1]))?
} else {
vm_port
};
(host_ip, host_port, vm_port)
}
_ => return Err("Expected format: [[hostip:][hostport]:]vmport".to_string()),
};
Ok(Self {
host_ip,
host_port,
vm_port,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct BindMount {
pub source: PathBuf,
pub dest: PathBuf,
pub read_only: bool,
}
impl BindMount {
pub fn tag(&self) -> String {
escape_path(&self.dest.to_string_lossy())
}
pub fn socket_name(&self) -> String {
format!("{}.sock", self.tag())
}
}
impl Display for BindMount {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let source = self.source.to_string_lossy();
let dest = self.dest.to_string_lossy();
if self.read_only {
write!(f, "{source}:{dest}:ro")
} else {
write!(f, "{source}:{dest}")
}
}
}
impl FromStr for BindMount {
type Err = String;
fn from_str(src: &str) -> Result<Self, Self::Err> {
let parts: Vec<&str> = src.split(':').collect();
if parts.len() != 2 && parts.len() != 3 {
return Err("Expected format: source:dest[:ro]".to_string());
}
let source = PathBuf::from(parts[0]);
if !source.is_absolute() {
return Err("source must be an absolute path".to_string());
}
if !source.is_dir() {
return Err("source doesn't exist or isn't a directory".to_string());
}
let dest = PathBuf::from(parts[1]);
if !dest.is_absolute() {
return Err("dest must be an absolute path".to_string());
}
if parts.len() == 3 {
let options = parts[2];
if options == "ro" {
return Ok(BindMount {
source,
dest,
read_only: true,
});
} else {
return Err("Expected format: source:dest[:ro]".to_string());
}
}
Ok(BindMount {
source,
dest,
read_only: false,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct PmemMount {
pub dest: PathBuf,
pub size: u64,
}
impl FromStr for PmemMount {
type Err = String;
fn from_str(src: &str) -> Result<Self, Self::Err> {
let parts: Vec<&str> = src.split(':').collect();
if parts.len() != 2 {
return Err("Expected format: dest:<size>".to_string());
}
let dest = PathBuf::from(parts[0]);
if !dest.is_absolute() {
return Err("dest must be an absolute path".to_string());
}
let size = if let Ok(size) = parts[1].parse() {
size
} else {
return Err("Couldn't parse size as integer".to_string());
};
Ok(PmemMount { dest, size })
}
}
#[derive(Clone, Debug, PartialEq, ValueEnum)]
pub enum Pull {
Missing,
Never,
Newer,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct EnvVar {
pub key: String,
pub value: String,
}
impl FromStr for EnvVar {
type Err = String;
fn from_str(src: &str) -> Result<Self, Self::Err> {
let parts: Vec<&str> = src.split('=').collect();
if parts.len() != 2 {
return Err("Expected format: KEY=VALUE".to_string());
}
Ok(Self {
key: parts[0].to_string(),
value: parts[1].to_string(),
})
}
}
#[derive(Debug, Clone, Subcommand)]
pub enum Command {
Ps(PsCommand),
Stop(StopCommand),
Exec(ExecCommand),
Run(RunCommand),
Ksm(KsmCommand),
Clean(CleanCommand),
Completions { shell: clap_complete::Shell },
Manpage { out_dir: PathBuf },
}
#[derive(Debug, Clone, Args)]
pub struct PsCommand {}
#[derive(Debug, Clone, Args)]
pub struct StopCommand {
#[arg(value_parser = parse_valid_vmid)]
pub vmid: String,
}
#[derive(Debug, Clone, Args)]
pub struct ExecCommand {
#[arg(short, long, value_parser = EnvVar::from_str)]
pub env: Vec<EnvVar>,
#[arg(
short,
long,
default_value = "20",
value_parser = parse_seconds_to_duration,
)]
pub ssh_timeout: Duration,
#[arg(short, long, default_value = "auto")]
pub interactive: Interactive,
#[arg(short, long, default_value = "auto")]
pub tty: Tty,
#[arg(value_parser = parse_valid_vmid)]
pub vmid: String,
pub args: Vec<String>,
}
#[derive(Debug, Clone, Args)]
pub struct RunCommand {
#[arg(short, long, conflicts_with_all = ["interactive", "tty"])]
pub detach: bool,
#[arg(long)]
pub rm: bool,
#[arg(long)]
pub disable_kvm: bool,
#[arg(short, long, value_parser = EnvVar::from_str)]
pub env: Vec<EnvVar>,
#[arg(short, long = "volume", value_parser= BindMount::from_str)]
pub volumes: Vec<BindMount>,
#[arg(long = "pmem", value_parser = PmemMount::from_str)]
pub pmems: Vec<PmemMount>,
#[arg(short, long = "publish", value_parser = PublishPort::from_str)]
pub published_ports: Vec<PublishPort>,
#[arg(
short,
long,
default_value = "20",
value_parser = parse_seconds_to_duration,
)]
pub ssh_timeout: Duration,
#[arg(short, long, default_value = "auto")]
pub interactive: Interactive,
#[arg(short, long, default_value = "auto")]
pub tty: Tty,
#[arg(long)]
pub show_vm_window: bool,
#[arg(long, default_value = "missing")]
pub pull: Pull,
#[arg(value_parser = OsTypeOrImagePath::from_str)]
pub image_source: OsTypeOrImagePath,
pub args: Vec<String>,
}
#[derive(Debug, Clone, Args)]
pub struct KsmEnableDisable {
#[arg(short, long)]
pub enable: bool,
#[arg(short, long)]
pub disable: bool,
}
#[derive(Debug, Clone, Args)]
pub struct KsmCommand {
#[command(flatten)]
pub ksm_enable_disable: Option<KsmEnableDisable>,
}
#[derive(Debug, Clone, Args)]
pub struct CleanCommand {}
#[derive(Debug, Clone, Parser)]
#[command(name = "vmexec", author, about, version)]
pub struct Cli {
#[arg(long, default_value = "warn")]
pub log_level: Level,
#[clap(subcommand)]
pub command: Command,
}
#[cfg(test)]
mod tests {
use super::*;
use pretty_assertions::assert_eq;
use rstest::rstest;
#[rstest]
#[case("archlinux", OsTypeOrImagePath::OsType(OsType::Archlinux))]
#[case(env!("CARGO_MANIFEST_PATH"), OsTypeOrImagePath::ImagePath(PathBuf::from(env!("CARGO_MANIFEST_PATH"))))]
fn test_parse_os_type_or_image_path(#[case] input: &str, #[case] expected: OsTypeOrImagePath) {
let actual = OsTypeOrImagePath::from_str(input).unwrap();
assert_eq!(actual, expected);
}
#[rstest]
#[case("just something", "Could not parse")]
fn test_parse_os_type_or_image_path_invalid(#[case] input: &str, #[case] expected: &str) {
let actual = OsTypeOrImagePath::from_str(input).unwrap_err();
assert!(actual.starts_with(expected));
}
#[rstest]
#[case("127.0.0.1:8080:80", "127.0.0.1", "8080", "80")]
#[case("80", "0.0.0.0", "80", "80")]
#[case("8080:80", "0.0.0.0", "8080", "80")]
#[case("127.0.0.1::80", "127.0.0.1", "80", "80")]
fn test_parse_publish_port_valid(
#[case] input: &str,
#[case] host_ip: Ipv4Addr,
#[case] host_port: u32,
#[case] vm_port: u32,
) {
let actual = PublishPort::from_str(input).unwrap();
let expected = PublishPort {
host_ip,
host_port,
vm_port,
};
assert_eq!(actual, expected);
}
#[rstest]
#[case("foo", "'foo' is not a valid port")]
#[case("foo::", "'foo' is not a valid IPv4")]
#[case("::", "Expected format: [[hostip:][hostport]:]vmport")]
#[case("1:2:3:4", "Expected format: [[hostip:][hostport]:]vmport")]
#[case(":80:", "Expected format: [[hostip:][hostport]:]vmport")]
fn test_parse_publish_port_invalid(#[case] input: &str, #[case] expected: &str) {
let actual = PublishPort::from_str(input).unwrap_err();
assert_eq!(actual, expected);
}
#[rstest]
#[case("/tmp:/tmp", "/tmp", "/tmp", false)]
#[case("/usr/bin:/somewhere/else", "/usr/bin", "/somewhere/else", false)]
#[case("/usr/bin:/somewhere/else:ro", "/usr/bin", "/somewhere/else", true)]
fn test_parse_bind_volume_valid(
#[case] input: &str,
#[case] source: PathBuf,
#[case] dest: PathBuf,
#[case] read_only: bool,
) {
let actual = BindMount::from_str(input).unwrap();
let expected = BindMount {
source,
dest,
read_only,
};
assert_eq!(actual, expected);
}
#[rstest]
#[case("tmp:/tmp", "source must be an absolute path")]
#[case("/nowhere:/tmp", "source doesn't exist or isn't a directory")]
#[case("/tmp:tmp", "dest must be an absolute path")]
#[case("/tmp", "Expected format: source:dest[:ro]")]
#[case("/tmp:/tmp:something", "Expected format: source:dest[:ro]")]
fn test_parse_bind_volume_invalid(#[case] input: &str, #[case] expected: &str) {
let actual = BindMount::from_str(input).unwrap_err();
assert_eq!(actual, expected);
}
#[rstest]
#[case("/tmp:2", "/tmp", 2)]
#[case("/tmp:200", "/tmp", 200)]
fn test_parse_pmem_valid(#[case] input: &str, #[case] dest: PathBuf, #[case] size: u64) {
let actual = PmemMount::from_str(input).unwrap();
let expected = PmemMount { dest, size };
assert_eq!(actual, expected);
}
#[rstest]
#[case("tmp:2", "dest must be an absolute path")]
fn test_parse_pmem_invalid(#[case] input: &str, #[case] expected: &str) {
let actual = PmemMount::from_str(input).unwrap_err();
assert_eq!(actual, expected);
}
#[rstest]
#[case("key=value", "key", "value")]
#[case("KEY=VALUE", "KEY", "VALUE")]
fn test_parse_env_var_valid(#[case] input: &str, #[case] key: String, #[case] value: String) {
let actual = EnvVar::from_str(input).unwrap();
let expected = EnvVar { key, value };
assert_eq!(actual, expected);
}
#[rstest]
#[case("keyvalue", "Expected format: KEY=VALUE")]
#[case("=key=value", "Expected format: KEY=VALUE")]
#[case("key=value=", "Expected format: KEY=VALUE")]
fn test_parse_env_var_invalid(#[case] input: &str, #[case] expected: &str) {
let actual = EnvVar::from_str(input).unwrap_err();
assert_eq!(actual, expected);
}
}