use std::io::Read;
use std::path::Path;
use std::process::Command;
use ssh2::Session;
use crate::config::{GuestOs, Result, TestbedError, VmProfile};
pub mod streaming;
pub struct VmSession {
pub session: Session,
pub port: u16,
pub os: GuestOs,
pub user: String,
}
pub fn connect(profile: &VmProfile) -> Result<VmSession> {
let port = profile.ssh_port;
connect_raw(port, profile.user, profile.os)
}
pub fn connect_from_port(port: u16, user: &str, os: GuestOs) -> Result<VmSession> {
connect_raw(port, user, os)
}
fn connect_raw(port: u16, user: &str, os: GuestOs) -> Result<VmSession> {
let tcp = std::net::TcpStream::connect(("127.0.0.1", port))
.map_err(|e| TestbedError::SshFailed {
port,
source: anyhow::anyhow!(e),
})?;
let mut session = Session::new().map_err(|e| TestbedError::SshFailed {
port,
source: anyhow::anyhow!(e),
})?;
session.set_tcp_stream(tcp);
session.handshake().map_err(|e| TestbedError::SshFailed {
port,
source: anyhow::anyhow!(e),
})?;
authenticate_raw(&mut session, user).map_err(|e| TestbedError::SshFailed {
port,
source: e,
})?;
Ok(VmSession {
session,
port,
os,
user: user.to_string(),
})
}
fn authenticate_raw(session: &mut Session, user: &str) -> anyhow::Result<()> {
if session.userauth_agent(user).is_ok() {
return Ok(());
}
let key_names = ["id_ed25519", "id_rsa", "id_ecdsa"];
if let Some(home) = dirs::home_dir() {
let ssh_dir = home.join(".ssh");
for key_name in &key_names {
let key_path = ssh_dir.join(key_name);
if key_path.exists()
&& session
.userauth_pubkey_file(user, None, &key_path, None)
.is_ok()
{
return Ok(());
}
}
}
if let Some(config_dir) = dirs::config_dir() {
let vagrant_key = config_dir.join("foundation_testbed/vagrant_insecure_key");
if vagrant_key.exists()
&& session
.userauth_pubkey_file(user, None, &vagrant_key, None)
.is_ok()
{
return Ok(());
}
}
anyhow::bail!("all auth methods failed for user '{user}'")
}
pub fn exec(session: &mut VmSession, cmd: &str) -> Result<String> {
let (output, _code) = exec_with_exit(session, cmd)?;
Ok(output)
}
pub fn exec_with_exit(session: &mut VmSession, cmd: &str) -> Result<(String, i32)> {
let wrapped = wrap_command(cmd, session.os);
let mut channel = session
.session
.channel_session()
.map_err(|e| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!(e),
})?;
channel
.exec(&wrapped)
.map_err(|e| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!(e),
})?;
let mut output = String::new();
channel
.read_to_string(&mut output)
.map_err(|e| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!(e),
})?;
let exit_code = channel
.exit_status()
.map_err(|e| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!(e),
})?;
Ok((output, exit_code))
}
pub fn exec_ps_windows(session: &mut VmSession, script: &str) -> Result<(String, i32)> {
let utf16le: Vec<u8> = script
.encode_utf16()
.flat_map(|c| c.to_le_bytes())
.collect();
let encoded =
base64::engine::Engine::encode(&base64::engine::general_purpose::STANDARD, &utf16le);
let cmd = format!("powershell -NoProfile -EncodedCommand {encoded}");
let (mut stdout, exit_code) = exec_with_exit(session, &cmd)?;
if let Some(idx) = stdout.find("#< CLIXML") {
stdout.truncate(idx);
}
Ok((stdout, exit_code))
}
pub fn upload(session: &mut VmSession, local: &Path, remote: &str) -> Result<()> {
let key_path = find_ssh_key().ok_or_else(|| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!("no SSH key found for SCP upload"),
})?;
let status = Command::new("scp")
.args([
"-P", &session.port.to_string(),
"-o", "StrictHostKeyChecking=no",
"-o", "UserKnownHostsFile=/dev/null",
"-o", "LogLevel=quiet",
"-i", &key_path,
local.to_str().ok_or_else(|| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!("local path {local:?} is not valid UTF-8"),
})?,
&format!("{}@127.0.0.1:{remote}", session.user),
])
.status()
.map_err(|e| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!("spawning scp: {e}"),
})?;
if !status.success() {
return Err(TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!("scp upload failed (exit {:?})", status.code()),
});
}
Ok(())
}
fn find_ssh_key() -> Option<String> {
if let Some(config_dir) = dirs::config_dir() {
let vagrant_key = config_dir.join("foundation_testbed/vagrant_insecure_key");
if vagrant_key.exists() {
return Some(vagrant_key.to_string_lossy().to_string());
}
}
let key_names = ["id_ed25519", "id_rsa", "id_ecdsa"];
if let Some(home) = dirs::home_dir() {
let ssh_dir = home.join(".ssh");
for key_name in &key_names {
let key_path = ssh_dir.join(key_name);
if key_path.exists() {
return Some(key_path.to_string_lossy().to_string());
}
}
}
None
}
pub fn download(session: &mut VmSession, remote: &str, local: &Path) -> Result<()> {
let key_path = find_ssh_key().ok_or_else(|| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!("no SSH key found for SCP download"),
})?;
let status = Command::new("scp")
.args([
"-P", &session.port.to_string(),
"-o", "StrictHostKeyChecking=no",
"-o", "UserKnownHostsFile=/dev/null",
"-o", "LogLevel=quiet",
"-i", &key_path,
&format!("{}@127.0.0.1:{remote}", session.user),
local.to_str().ok_or_else(|| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!("local path {local:?} is not valid UTF-8"),
})?,
])
.status()
.map_err(|e| TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!("spawning scp: {e}"),
})?;
if !status.success() {
return Err(TestbedError::SshFailed {
port: session.port,
source: anyhow::anyhow!("scp download failed (exit {:?})", status.code()),
});
}
Ok(())
}
pub fn check(profile: &VmProfile) -> Result<()> {
let mut session = connect(profile)?;
let output = exec(&mut session, "echo testbed-ping")?;
if output.trim() != "testbed-ping" {
return Err(TestbedError::SshFailed {
port: profile.ssh_port,
source: anyhow::anyhow!("echo test returned unexpected output: {output:?}"),
});
}
Ok(())
}
pub fn shell(profile: &VmProfile) -> Result<()> {
let mut args = vec![
"-p".to_string(),
profile.ssh_port.to_string(),
"-o".to_string(),
"StrictHostKeyChecking=no".to_string(),
"-o".to_string(),
"UserKnownHostsFile=/dev/null".to_string(),
"-o".to_string(),
"LogLevel=quiet".to_string(),
format!("{}@127.0.0.1", profile.user),
];
if profile.os == GuestOs::Linux {
args.insert(0, "-tt".to_string());
}
let status = Command::new("ssh")
.args(&args)
.status()
.map_err(|e| TestbedError::SshFailed {
port: profile.ssh_port,
source: anyhow::anyhow!("spawning ssh: {e}"),
})?;
if !status.success() {
return Err(TestbedError::SshFailed {
port: profile.ssh_port,
source: anyhow::anyhow!("ssh exited with status {status:?}"),
});
}
Ok(())
}
pub fn wrap_command(cmd: &str, os: GuestOs) -> String {
match os {
GuestOs::Linux | GuestOs::MacOS => wrap_command_bash(cmd),
GuestOs::Windows => wrap_command_cmd(cmd),
}
}
pub fn wrap_command_cmd(cmd: &str) -> String {
format!("cmd /c {cmd:?}")
}
pub fn wrap_command_bash(cmd: &str) -> String {
format!("bash -c {cmd:?}")
}