use std::collections::HashMap;
use std::fmt;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use std::time::Duration;
use bevy_app::{App, First, Last, PostUpdate};
use bevy_ecs::message::Message;
use bevy_ecs::resource::Resource;
use bevy_ecs::schedule::common_conditions::on_message;
use bevy_ecs::schedule::IntoScheduleConfigs;
use crate::credentials::Secret;
use crate::request::RequestId;
use crate::response::BackendError;
use crate::BackendSystems;
mod known_hosts;
mod russh_client;
#[cfg(feature = "sftp")]
mod sftp;
mod ssh_config;
mod systems;
mod transport;
pub use russh_client::RusshTransport;
#[cfg(feature = "sftp")]
pub use sftp::{SftpEntry, SftpEntryKind, SftpFinished, SftpOp, SftpOutcome, SftpProgress};
pub use transport::{FakeSshTransport, SshConnId, SshEvent, SshTransport, SshTransportRes};
pub const DEFAULT_SSH_CONNECT_TIMEOUT: Duration = Duration::from_secs(15);
pub const DEFAULT_SSH_COMMAND_TIMEOUT: Duration = Duration::from_secs(60);
pub const DEFAULT_SSH_MAX_OUTPUT_BYTES: u64 = 8 * 1024 * 1024;
pub const DEFAULT_SFTP_TIMEOUT: Duration = Duration::from_secs(300);
pub const DEFAULT_SFTP_MAX_BYTES: u64 = 256 * 1024 * 1024;
pub const MAX_SSH_COMMAND_BYTES: usize = 64 * 1024;
#[derive(Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct SshName(Arc<str>);
impl SshName {
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Debug for SshName {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "SshName({:?})", &*self.0)
}
}
impl fmt::Display for SshName {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl From<&str> for SshName {
fn from(name: &str) -> Self {
Self(Arc::from(name))
}
}
impl From<String> for SshName {
fn from(name: String) -> Self {
Self(Arc::from(name))
}
}
impl From<&String> for SshName {
fn from(name: &String) -> Self {
Self(Arc::from(name.as_str()))
}
}
impl From<&SshName> for SshName {
fn from(name: &SshName) -> Self {
name.clone()
}
}
impl PartialEq<str> for SshName {
fn eq(&self, other: &str) -> bool {
&*self.0 == other
}
}
impl PartialEq<&str> for SshName {
fn eq(&self, other: &&str) -> bool {
&*self.0 == *other
}
}
#[derive(Clone, Debug)]
pub struct SshSettings {
pub(crate) allow_in_release: bool,
pub(crate) max_connections: usize,
pub(crate) max_requests: usize,
}
impl Default for SshSettings {
fn default() -> Self {
Self { allow_in_release: false, max_connections: 16, max_requests: 256 }
}
}
impl SshSettings {
pub fn allow_in_release(mut self, allow: bool) -> Self {
self.allow_in_release = allow;
self
}
pub fn with_max_connections(mut self, connections: usize) -> Self {
self.max_connections = connections.max(1);
self
}
pub fn with_max_requests_per_connection(mut self, requests: usize) -> Self {
self.max_requests = requests.max(1);
self
}
pub fn is_allowed(&self) -> bool {
allowed(cfg!(debug_assertions), self.allow_in_release)
}
}
pub(crate) fn allowed(debug_build: bool, allow_in_release: bool) -> bool {
debug_build || allow_in_release
}
#[derive(Clone)]
pub struct SshAuth(pub(crate) AuthKind);
#[derive(Clone)]
pub(crate) enum AuthKind {
KeyFile { path: PathBuf, passphrase: Option<Secret> },
Agent,
Password(Secret),
KeyboardInteractive(Arc<dyn SshPromptResponder>),
}
impl SshAuth {
pub fn key_file(path: impl Into<PathBuf>) -> Self {
Self(AuthKind::KeyFile { path: path.into(), passphrase: None })
}
pub fn key_file_with_passphrase(path: impl Into<PathBuf>, passphrase: impl Into<Secret>) -> Self {
Self(AuthKind::KeyFile { path: path.into(), passphrase: Some(passphrase.into()) })
}
pub fn agent() -> Self {
Self(AuthKind::Agent)
}
pub fn password(password: impl Into<Secret>) -> Self {
Self(AuthKind::Password(password.into()))
}
pub fn keyboard_interactive(responder: impl SshPromptResponder) -> Self {
Self(AuthKind::KeyboardInteractive(Arc::new(responder)))
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct SshPrompt {
pub text: String,
pub echo: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct SshPromptRequest {
pub name: String,
pub instructions: String,
pub prompts: Vec<SshPrompt>,
}
impl SshPromptRequest {
pub fn new(prompts: Vec<SshPrompt>) -> Self {
Self { name: String::new(), instructions: String::new(), prompts }
}
}
impl SshPrompt {
pub fn new(text: impl Into<String>, echo: bool) -> Self {
Self { text: text.into(), echo }
}
}
pub trait SshPromptResponder: Send + Sync + 'static {
fn respond(&self, request: &SshPromptRequest) -> Option<Vec<Secret>>;
}
#[derive(Clone, Default)]
pub struct SshPromptAnswers {
answers: Vec<(String, Secret)>,
}
impl fmt::Debug for SshPromptAnswers {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshPromptAnswers").field("words", &self.answers.iter().map(|(w, _)| w.as_str()).collect::<Vec<_>>()).finish()
}
}
impl SshPromptAnswers {
pub fn new() -> Self {
Self::default()
}
pub fn answer_containing(mut self, word: impl Into<String>, answer: impl Into<Secret>) -> Self {
self.answers.push((word.into().to_lowercase(), answer.into()));
self
}
}
impl SshPromptResponder for SshPromptAnswers {
fn respond(&self, request: &SshPromptRequest) -> Option<Vec<Secret>> {
request
.prompts
.iter()
.map(|prompt| {
let text = prompt.text.to_lowercase();
self.answers.iter().find(|(word, _)| text.contains(word.as_str())).map(|(_, answer)| answer.clone())
})
.collect()
}
}
#[derive(Clone, Debug)]
pub struct SshReconnect {
pub(crate) base: Duration,
pub(crate) cap: Duration,
pub(crate) max_attempts: Option<u32>,
pub(crate) stable_after: Duration,
pub(crate) jitter: bool,
}
impl Default for SshReconnect {
fn default() -> Self {
Self { base: Duration::from_secs(1), cap: Duration::from_secs(30), max_attempts: None, stable_after: Duration::from_secs(10), jitter: true }
}
}
impl SshReconnect {
pub fn with_base(mut self, base: Duration) -> Self {
self.base = base.clamp(Duration::from_millis(1), crate::config::MAX_TIMEOUT);
self
}
pub fn with_cap(mut self, cap: Duration) -> Self {
self.cap = cap.min(crate::config::MAX_TIMEOUT);
self
}
pub fn with_max_attempts(mut self, max: Option<u32>) -> Self {
self.max_attempts = max;
self
}
pub fn with_stable_after(mut self, stable_after: Duration) -> Self {
self.stable_after = stable_after.min(crate::config::MAX_TIMEOUT);
self
}
pub fn with_jitter(mut self, jitter: bool) -> Self {
self.jitter = jitter;
self
}
pub fn delay_bound(&self, attempt: u32) -> Duration {
let factor = 2u32.checked_pow(attempt.saturating_sub(1).min(30)).unwrap_or(u32::MAX);
self.base.saturating_mul(factor).min(self.cap.max(self.base))
}
pub(crate) fn delay(&self, attempt: u32, random: u64) -> Duration {
let bound = self.delay_bound(attempt);
if !self.jitter {
return bound;
}
let nanos = u64::try_from(bound.as_nanos()).unwrap_or(u64::MAX);
Duration::from_nanos(random % nanos.saturating_add(1))
}
pub(crate) fn may_retry(&self, attempt: u32) -> bool {
self.max_attempts.is_none_or(|max| attempt <= max)
}
}
pub(crate) fn file_name(path: &Path) -> String {
path.file_name().map_or_else(|| "<no file name>".to_string(), |n| n.to_string_lossy().into_owned())
}
impl fmt::Debug for SshAuth {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.0 {
AuthKind::KeyFile { path, passphrase } => {
f.debug_struct("KeyFile").field("file", &file_name(path)).field("passphrase", &passphrase.as_ref().map(|_| "<redacted>")).finish()
}
AuthKind::Agent => f.write_str("Agent"),
AuthKind::Password(_) => f.write_str("Password(<redacted>)"),
AuthKind::KeyboardInteractive(_) => f.write_str("KeyboardInteractive"),
}
}
}
#[derive(Clone)]
pub(crate) enum TargetSource {
Direct,
Config(Option<PathBuf>),
}
#[derive(Clone)]
pub struct SshTarget {
pub(crate) host: String,
pub(crate) source: TargetSource,
pub(crate) port: Option<u16>,
pub(crate) user: Option<String>,
pub(crate) auth: Vec<SshAuth>,
pub(crate) known_hosts: Vec<PathBuf>,
pub(crate) pinned: Vec<String>,
pub(crate) connect_timeout: Duration,
pub(crate) keepalive_interval: Duration,
pub(crate) keepalive_max: u32,
pub(crate) command_timeout: Duration,
pub(crate) max_output_bytes: u64,
pub(crate) sftp_timeout: Duration,
pub(crate) max_transfer_bytes: u64,
pub(crate) max_channels: usize,
pub(crate) reconnect: Option<SshReconnect>,
pub(crate) allow_terrapin_vulnerable: bool,
}
impl fmt::Debug for SshTarget {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshTarget")
.field("host", &self.host)
.field("from_ssh_config", &matches!(self.source, TargetSource::Config(_)))
.field("port", &self.port)
.field("user", &self.user)
.field("auth", &self.auth)
.field("known_hosts_files", &self.known_hosts.len())
.field("pinned_fingerprints", &self.pinned.len())
.field("connect_timeout", &self.connect_timeout)
.field("keepalive", &(self.keepalive_interval, self.keepalive_max))
.field("command_timeout", &self.command_timeout)
.field("max_output_bytes", &self.max_output_bytes)
.field("sftp_timeout", &self.sftp_timeout)
.field("max_transfer_bytes", &self.max_transfer_bytes)
.field("max_channels", &self.max_channels)
.field("reconnect", &self.reconnect)
.field("allow_terrapin_vulnerable", &self.allow_terrapin_vulnerable)
.finish()
}
}
impl SshTarget {
pub fn new(host: impl Into<String>, user: impl Into<String>) -> Self {
Self::blank(host.into(), TargetSource::Direct, Some(user.into()))
}
pub fn from_ssh_config(alias: impl Into<String>) -> Self {
Self::blank(alias.into(), TargetSource::Config(None), None)
}
pub fn from_ssh_config_file(path: impl Into<PathBuf>, alias: impl Into<String>) -> Self {
Self::blank(alias.into(), TargetSource::Config(Some(path.into())), None)
}
fn blank(host: String, source: TargetSource, user: Option<String>) -> Self {
Self {
host,
source,
port: None,
user,
auth: Vec::new(),
known_hosts: Vec::new(),
pinned: Vec::new(),
connect_timeout: DEFAULT_SSH_CONNECT_TIMEOUT,
keepalive_interval: Duration::from_secs(15),
keepalive_max: 3,
command_timeout: DEFAULT_SSH_COMMAND_TIMEOUT,
max_output_bytes: DEFAULT_SSH_MAX_OUTPUT_BYTES,
sftp_timeout: DEFAULT_SFTP_TIMEOUT,
max_transfer_bytes: DEFAULT_SFTP_MAX_BYTES,
max_channels: 8,
reconnect: None,
allow_terrapin_vulnerable: false,
}
}
pub fn with_port(mut self, port: u16) -> Self {
self.port = Some(port);
self
}
pub fn with_user(mut self, user: impl Into<String>) -> Self {
self.user = Some(user.into());
self
}
pub fn with_auth(mut self, auth: SshAuth) -> Self {
self.auth.push(auth);
self
}
pub fn with_known_hosts_file(mut self, path: impl Into<PathBuf>) -> Self {
self.known_hosts.push(path.into());
self
}
pub fn trust_host_key_fingerprint(mut self, fingerprint: impl Into<String>) -> Self {
self.pinned.push(fingerprint.into().trim().to_string());
self
}
pub fn with_connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = timeout.clamp(Duration::from_secs(1), crate::config::MAX_TIMEOUT);
self
}
pub fn with_keepalive(mut self, interval: Duration, max_missed: u32) -> Self {
self.keepalive_interval = interval.clamp(Duration::from_secs(1), crate::config::MAX_TIMEOUT);
self.keepalive_max = max_missed.clamp(1, 100);
self
}
pub fn with_command_timeout(mut self, timeout: Duration) -> Self {
self.command_timeout = timeout.clamp(Duration::from_millis(1), crate::config::MAX_TIMEOUT);
self
}
pub fn with_max_output_bytes(mut self, bytes: u64) -> Self {
self.max_output_bytes = bytes.max(1024);
self
}
pub fn with_sftp_timeout(mut self, timeout: Duration) -> Self {
self.sftp_timeout = timeout.clamp(Duration::from_millis(1), crate::config::MAX_TIMEOUT);
self
}
pub fn with_max_transfer_bytes(mut self, bytes: u64) -> Self {
self.max_transfer_bytes = bytes.max(1024);
self
}
pub fn with_max_channels(mut self, channels: usize) -> Self {
self.max_channels = channels.clamp(1, 64);
self
}
pub fn with_reconnect(mut self, reconnect: SshReconnect) -> Self {
self.reconnect = Some(reconnect);
self
}
pub fn allow_terrapin_vulnerable(mut self, allow: bool) -> Self {
self.allow_terrapin_vulnerable = allow;
self
}
pub fn host(&self) -> &str {
&self.host
}
pub fn port(&self) -> Option<u16> {
self.port
}
pub fn user(&self) -> Option<&str> {
self.user.as_deref()
}
pub fn connect_timeout(&self) -> Duration {
self.connect_timeout
}
pub fn command_timeout(&self) -> Duration {
self.command_timeout
}
pub fn max_output_bytes(&self) -> u64 {
self.max_output_bytes
}
pub fn sftp_timeout(&self) -> Duration {
self.sftp_timeout
}
pub fn max_transfer_bytes(&self) -> u64 {
self.max_transfer_bytes
}
pub fn validate(&self) -> Result<(), BackendError> {
check_host(&self.host)?;
if let Some(user) = &self.user {
check_user(user)?;
}
if self.port == Some(0) {
return Err(BackendError::InvalidRequest("the SSH port must not be 0".into()));
}
for pin in &self.pinned {
check_fingerprint(pin)?;
}
match self.source {
TargetSource::Direct if self.user.is_none() => Err(BackendError::InvalidRequest("no SSH user given".into())),
TargetSource::Direct if self.auth.is_empty() => {
Err(BackendError::InvalidRequest("no SSH authentication method given (SshAuth::key_file or SshAuth::agent)".into()))
}
_ => Ok(()),
}
}
}
pub(crate) fn check_host(host: &str) -> Result<(), BackendError> {
if host.is_empty() || host.len() > 255 || host.starts_with('-') || host.chars().any(|c| c.is_whitespace() || c.is_control()) {
return Err(BackendError::InvalidRequest("the SSH host is not a valid host name or address".into()));
}
Ok(())
}
pub(crate) fn check_user(user: &str) -> Result<(), BackendError> {
if user.is_empty() || user.len() > 255 || user.chars().any(|c| c.is_whitespace() || c.is_control()) {
return Err(BackendError::InvalidRequest("the SSH user name is empty or has whitespace / control characters".into()));
}
Ok(())
}
pub(crate) fn check_fingerprint(pin: &str) -> Result<(), BackendError> {
let ok = pin.strip_prefix("SHA256:").is_some_and(|b64| b64.len() == 43 && b64.bytes().all(|b| b.is_ascii_alphanumeric() || b == b'+' || b == b'/'));
if ok {
Ok(())
} else {
Err(BackendError::InvalidRequest(
"a pinned host key fingerprint must look like `SHA256:` + 43 base64 characters (as `ssh-keygen -lf` prints it)".into(),
))
}
}
#[derive(Clone)]
pub struct SshCommand {
pub(crate) command: String,
pub(crate) timeout: Option<Duration>,
pub(crate) max_output_bytes: Option<u64>,
pub(crate) stdin: Option<Vec<u8>>,
}
impl fmt::Debug for SshCommand {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshCommand")
.field("command_bytes", &self.command.len())
.field("timeout", &self.timeout)
.field("max_output_bytes", &self.max_output_bytes)
.field("stdin_bytes", &self.stdin.as_ref().map(Vec::len))
.finish()
}
}
impl SshCommand {
pub fn new(command: impl Into<String>) -> Self {
Self { command: command.into(), timeout: None, max_output_bytes: None, stdin: None }
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout.clamp(Duration::from_millis(1), crate::config::MAX_TIMEOUT));
self
}
pub fn with_max_output_bytes(mut self, bytes: u64) -> Self {
self.max_output_bytes = Some(bytes.max(1024));
self
}
pub fn with_stdin(mut self, stdin: impl Into<Vec<u8>>) -> Self {
self.stdin = Some(stdin.into());
self
}
pub fn command(&self) -> &str {
&self.command
}
pub fn timeout(&self) -> Option<Duration> {
self.timeout
}
pub fn max_output_bytes(&self) -> Option<u64> {
self.max_output_bytes
}
pub fn stdin(&self) -> Option<&[u8]> {
self.stdin.as_deref()
}
pub(crate) fn validate(&self) -> Result<(), BackendError> {
if self.command.trim().is_empty() {
return Err(BackendError::InvalidRequest("the SSH command is empty".into()));
}
if self.command.contains('\0') {
return Err(BackendError::InvalidRequest("the SSH command contains a NUL byte".into()));
}
if self.command.len() > MAX_SSH_COMMAND_BYTES {
return Err(BackendError::RequestTooLarge { limit: MAX_SSH_COMMAND_BYTES as u64, size: u64::try_from(self.command.len()).unwrap_or(u64::MAX) });
}
Ok(())
}
}
impl From<&str> for SshCommand {
fn from(command: &str) -> Self {
Self::new(command)
}
}
impl From<String> for SshCommand {
fn from(command: String) -> Self {
Self::new(command)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct SshExit {
pub status: Option<u32>,
pub signal: Option<String>,
pub stdout_bytes: u64,
pub stderr_bytes: u64,
}
impl SshExit {
pub fn with_status(status: u32) -> Self {
Self { status: Some(status), signal: None, stdout_bytes: 0, stderr_bytes: 0 }
}
pub fn with_signal(signal: impl Into<String>) -> Self {
Self { status: None, signal: Some(signal.into()), stdout_bytes: 0, stderr_bytes: 0 }
}
pub fn with_output_bytes(mut self, stdout: u64, stderr: u64) -> Self {
self.stdout_bytes = stdout;
self.stderr_bytes = stderr;
self
}
pub fn success(&self) -> bool {
self.status == Some(0)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum SshStream {
Stdout,
Stderr,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum SshState {
Connecting,
Connected,
Reconnecting {
attempt: u32,
retry_in: Duration,
},
Disconnected,
}
#[derive(Message, Clone, Debug)]
#[non_exhaustive]
pub struct SshStateChanged {
pub name: SshName,
pub state: SshState,
pub error: Option<BackendError>,
}
#[derive(Message, Clone)]
#[non_exhaustive]
pub struct SshOutput {
pub id: RequestId,
pub name: SshName,
pub stream: SshStream,
pub data: Vec<u8>,
}
impl SshOutput {
pub fn text(&self) -> String {
String::from_utf8_lossy(&self.data).into_owned()
}
}
impl fmt::Debug for SshOutput {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshOutput").field("id", &self.id).field("name", &self.name).field("stream", &self.stream).field("bytes", &self.data.len()).finish()
}
}
#[derive(Message, Clone, Debug)]
#[non_exhaustive]
pub struct SshFinished {
pub id: RequestId,
pub name: SshName,
pub started: Option<bool>,
pub result: Result<SshExit, BackendError>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct SshConnectionInfo {
pub state: SshState,
pub fingerprint: Option<String>,
pub last_error: Option<BackendError>,
pub pending_requests: usize,
pub attempt: u32,
}
#[derive(Resource, Default, Debug)]
pub struct SshConnections {
pub(crate) map: HashMap<SshName, SshConnectionInfo>,
}
impl SshConnections {
pub fn get(&self, name: &str) -> Option<&SshConnectionInfo> {
self.map.get(&SshName::from(name))
}
pub fn state(&self, name: &str) -> Option<SshState> {
self.get(name).map(|c| c.state)
}
pub fn is_connected(&self, name: &str) -> bool {
self.state(name) == Some(SshState::Connected)
}
pub fn iter(&self) -> impl Iterator<Item = (&SshName, &SshConnectionInfo)> {
self.map.iter()
}
}
pub(crate) enum SshQueued {
Connect(SshName, Box<SshTarget>),
Disconnect(SshName),
Run {
name: SshName,
id: RequestId,
command: SshCommand,
},
#[cfg(feature = "sftp")]
Sftp {
name: SshName,
id: RequestId,
op: SftpOp,
},
}
impl SshQueued {
pub(crate) fn request_id(&self) -> Option<RequestId> {
match self {
SshQueued::Run { id, .. } => Some(*id),
#[cfg(feature = "sftp")]
SshQueued::Sftp { id, .. } => Some(*id),
SshQueued::Connect(..) | SshQueued::Disconnect(_) => None,
}
}
}
#[derive(Resource, Default)]
pub struct SshClient {
queue: Mutex<Vec<SshQueued>>,
cancels: crate::inflight::CancelList,
}
impl fmt::Debug for SshClient {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshClient").field("queued", &self.lock().len()).finish_non_exhaustive()
}
}
impl SshClient {
fn lock(&self) -> MutexGuard<'_, Vec<SshQueued>> {
self.queue.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn drain(&self) -> Vec<SshQueued> {
std::mem::take(&mut *self.lock())
}
#[cfg(feature = "sftp")]
pub(crate) fn push_sftp(&self, name: SshName, op: SftpOp) -> RequestId {
let id = RequestId::next();
self.lock().push(SshQueued::Sftp { name, id, op });
id
}
pub fn connect(&self, name: impl Into<SshName>, target: SshTarget) {
self.lock().push(SshQueued::Connect(name.into(), Box::new(target)));
}
pub fn disconnect(&self, name: impl Into<SshName>) {
self.lock().push(SshQueued::Disconnect(name.into()));
}
pub fn run(&self, name: impl Into<SshName>, command: impl Into<SshCommand>) -> RequestId {
let id = RequestId::next();
self.lock().push(SshQueued::Run { name: name.into(), id, command: command.into() });
id
}
pub fn cancel(&self, id: RequestId) {
self.cancels.push(id);
}
pub(crate) fn share_cancels(&mut self, cancels: crate::inflight::CancelList) {
self.cancels = cancels;
}
}
pub(crate) fn build(app: &mut App, settings: &SshSettings) {
if settings.is_allowed() {
if cfg!(debug_assertions) {
tracing::info!(">>> NET-BACKEND: SSH available (feature `ssh`; admin / dev builds only)");
} else {
tracing::warn!(">>> NET-BACKEND: SSH ENABLED IN A RELEASE BUILD (allow_in_release): never ship SSH keys to players");
}
} else {
tracing::info!(">>> NET-BACKEND: SSH is compiled in but disabled in this release build; every SSH request is refused");
}
app.init_resource::<SshClient>();
let cancels = app.world().resource::<crate::InFlight>().cancel_list();
if let Some(mut client) = app.world_mut().get_resource_mut::<SshClient>() {
client.share_cancels(cancels);
}
app.insert_resource(systems::SshRuntime::new(settings.clone()))
.init_resource::<SshConnections>()
.add_message::<SshStateChanged>()
.add_message::<SshOutput>()
.add_message::<SshFinished>();
#[cfg(feature = "sftp")]
app.add_message::<SftpProgress>().add_message::<SftpFinished>();
#[cfg(feature = "ws")]
app.add_systems(First, systems::ssh_receive.in_set(BackendSystems::Receive).after(crate::ws::ws_receive))
.add_systems(PostUpdate, systems::ssh_send.in_set(BackendSystems::Send).after(crate::ws::ws_send))
.add_systems(Last, systems::ssh_exit.in_set(BackendSystems::Exit).after(crate::ws::ws_exit).run_if(on_message::<bevy_app::AppExit>));
#[cfg(not(feature = "ws"))]
app.add_systems(First, systems::ssh_receive.in_set(BackendSystems::Receive).after(crate::inflight::receive_answers))
.add_systems(PostUpdate, systems::ssh_send.in_set(BackendSystems::Send).after(crate::inflight::send_requests))
.add_systems(Last, systems::ssh_exit.in_set(BackendSystems::Exit).after(crate::inflight::shutdown_on_exit).run_if(on_message::<bevy_app::AppExit>));
if !app.world().contains_resource::<SshTransportRes>() {
app.insert_resource(SshTransportRes::new(RusshTransport::new().with_release_allowed(settings.allow_in_release)));
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::{allowed, SshAuth, SshCommand, SshPrompt, SshPromptAnswers, SshPromptRequest, SshPromptResponder, SshReconnect, SshSettings, SshTarget};
use crate::BackendError;
#[test]
fn prompt_answers_match_words_and_give_up_on_unknown_prompts() {
let answers = SshPromptAnswers::new().answer_containing("Password", "fake-pw-1").answer_containing("code", "123456");
let request = SshPromptRequest::new(vec![SshPrompt::new("Password: ", false), SshPrompt::new("Verification code: ", true)]);
let got: Option<Vec<String>> = answers.respond(&request).map(|a| a.iter().map(|s| s.expose().to_string()).collect());
assert_eq!(got, Some(vec!["fake-pw-1".to_string(), "123456".to_string()]));
assert!(answers.respond(&SshPromptRequest::new(vec![SshPrompt::new("Favourite colour?", true)])).is_none());
assert_eq!(answers.respond(&SshPromptRequest::new(Vec::new())).map(|a| a.len()), Some(0));
let debug = format!("{answers:?} {:?} {:?}", SshAuth::password("fake-pw-1"), SshAuth::keyboard_interactive(answers.clone()));
assert!(!debug.contains("fake-pw-1") && !debug.contains("123456"), "{debug}");
}
#[test]
fn reconnect_backoff_is_bounded() {
let policy = SshReconnect::default().with_base(Duration::from_millis(100)).with_cap(Duration::from_secs(2)).with_jitter(false);
assert_eq!(policy.delay_bound(1), Duration::from_millis(100));
assert_eq!(policy.delay_bound(3), Duration::from_millis(400));
assert_eq!(policy.delay_bound(40), Duration::from_secs(2));
assert_eq!(policy.delay(u32::MAX, u64::MAX), Duration::from_secs(2));
let jitter = SshReconnect::default().with_base(Duration::MAX).with_cap(Duration::MAX);
assert!(jitter.delay(5, u64::MAX) <= crate::MAX_TIMEOUT);
assert!(SshReconnect::default().with_max_attempts(Some(2)).may_retry(2));
assert!(!SshReconnect::default().with_max_attempts(Some(2)).may_retry(3));
}
#[test]
fn the_release_guard() {
assert!(allowed(true, false), "debug builds always");
assert!(!allowed(false, false), "release builds refuse by default");
assert!(allowed(false, true), "release builds with allow_in_release");
assert_eq!(SshSettings::default().is_allowed(), cfg!(debug_assertions));
assert!(SshSettings::default().allow_in_release(true).is_allowed());
}
#[test]
fn targets_are_validated() {
let ok = SshTarget::new("host.example.com", "deploy").with_auth(SshAuth::agent());
assert!(ok.validate().is_ok());
assert!(ok.clone().with_port(0).validate().is_err());
assert!(SshTarget::new("host.example.com", "deploy").validate().is_err(), "no auth");
assert!(SshTarget::new("", "deploy").with_auth(SshAuth::agent()).validate().is_err());
assert!(SshTarget::new("-oProxyCommand=x", "deploy").with_auth(SshAuth::agent()).validate().is_err());
assert!(SshTarget::new("host", "de ploy").with_auth(SshAuth::agent()).validate().is_err());
assert!(SshTarget::new(
"host", "deploy
"
)
.with_auth(SshAuth::agent())
.validate()
.is_err());
let pin = format!("SHA256:{}", "A".repeat(43));
assert!(ok.clone().trust_host_key_fingerprint(pin).validate().is_ok());
assert!(ok.clone().trust_host_key_fingerprint("SHA256:short").validate().is_err());
assert!(ok.clone().trust_host_key_fingerprint("MD5:aa:bb").validate().is_err());
assert!(SshTarget::from_ssh_config("build").validate().is_ok());
}
#[test]
fn durations_and_limits_are_clamped() {
let target = SshTarget::new("h", "u")
.with_connect_timeout(Duration::MAX)
.with_command_timeout(Duration::MAX)
.with_sftp_timeout(Duration::ZERO)
.with_keepalive(Duration::MAX, 0)
.with_max_output_bytes(0)
.with_max_transfer_bytes(0)
.with_max_channels(10_000);
assert_eq!(target.connect_timeout(), crate::MAX_TIMEOUT);
assert_eq!(target.command_timeout(), crate::MAX_TIMEOUT);
assert_eq!(target.sftp_timeout(), Duration::from_millis(1));
assert_eq!((target.keepalive_interval, target.keepalive_max), (crate::MAX_TIMEOUT, 1));
assert_eq!((target.max_output_bytes(), target.max_transfer_bytes(), target.max_channels), (1024, 1024, 64));
let command = SshCommand::new("x").with_timeout(Duration::MAX).with_max_output_bytes(1);
assert_eq!((command.timeout(), command.max_output_bytes()), (Some(crate::MAX_TIMEOUT), Some(1024)));
}
#[test]
fn debug_output_never_shows_secrets_or_full_paths() {
let target = SshTarget::new("host", "deploy")
.with_auth(SshAuth::key_file_with_passphrase("/home/someone/.ssh/id_ed25519", "fake-pass-1234"))
.with_known_hosts_file("/home/someone/.ssh/known_hosts");
let debug = format!("{target:?}");
assert!(!debug.contains("fake-pass") && !debug.contains("someone"), "{debug}");
assert!(debug.contains("id_ed25519"));
let error = BackendError::AuthFailed("the server accepted none of: key file `id_ed25519`".into());
assert!(error.to_string().starts_with("SSH authentication failed"));
assert_eq!(error.was_sent(), Some(false));
}
}