use std::collections::{HashMap, HashSet, VecDeque};
use std::fmt;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use bevy_ecs::resource::Resource;
#[cfg(feature = "sftp")]
use super::{SftpOp, SftpOutcome};
use super::{SshCommand, SshExit, SshStream, SshTarget};
use crate::request::RequestId;
use crate::response::BackendError;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct SshConnId(u64);
static NEXT_CONN: AtomicU64 = AtomicU64::new(1);
impl SshConnId {
pub(crate) fn next() -> Self {
Self(NEXT_CONN.fetch_add(1, Ordering::Relaxed))
}
}
impl fmt::Display for SshConnId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "ssh#{}", self.0)
}
}
#[derive(Clone)]
#[non_exhaustive]
pub enum SshEvent {
Connected {
conn: SshConnId,
fingerprint: String,
},
Closed {
conn: SshConnId,
error: Option<BackendError>,
},
Started {
id: RequestId,
},
Output {
id: RequestId,
stream: SshStream,
data: Vec<u8>,
},
Finished {
id: RequestId,
result: Result<SshExit, BackendError>,
},
#[cfg(feature = "sftp")]
Progress {
id: RequestId,
done: u64,
total: Option<u64>,
},
#[cfg(feature = "sftp")]
SftpFinished {
id: RequestId,
result: Result<SftpOutcome, BackendError>,
},
}
impl fmt::Debug for SshEvent {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
SshEvent::Connected { conn, fingerprint } => f.debug_struct("Connected").field("conn", conn).field("fingerprint", fingerprint).finish(),
SshEvent::Closed { conn, error } => f.debug_struct("Closed").field("conn", conn).field("error", error).finish(),
SshEvent::Started { id } => f.debug_struct("Started").field("id", id).finish(),
SshEvent::Output { id, stream, data } => f.debug_struct("Output").field("id", id).field("stream", stream).field("bytes", &data.len()).finish(),
SshEvent::Finished { id, result } => f.debug_struct("Finished").field("id", id).field("result", result).finish(),
#[cfg(feature = "sftp")]
SshEvent::Progress { id, done, total } => f.debug_struct("Progress").field("id", id).field("done", done).field("total", total).finish(),
#[cfg(feature = "sftp")]
SshEvent::SftpFinished { id, result } => f.debug_struct("SftpFinished").field("id", id).field("result", result).finish(),
}
}
}
pub trait SshTransport: Send + Sync + 'static {
fn connect(&mut self, conn: SshConnId, target: SshTarget);
fn run(&mut self, conn: SshConnId, id: RequestId, command: SshCommand);
#[cfg(feature = "sftp")]
fn sftp(&mut self, conn: SshConnId, id: RequestId, op: SftpOp) {
let _ = (conn, id, op);
}
fn supports_sftp(&self) -> bool {
false
}
fn cancel(&mut self, id: RequestId) {
let _ = id;
}
fn close(&mut self, conn: SshConnId) {
let _ = conn;
}
fn poll(&mut self) -> Vec<SshEvent>;
fn shutdown(&mut self) {}
}
static NEXT_GENERATION: AtomicU64 = AtomicU64::new(1);
#[derive(Resource)]
pub struct SshTransportRes {
inner: Box<dyn SshTransport>,
generation: u64,
}
impl fmt::Debug for SshTransportRes {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshTransportRes").field("generation", &self.generation).finish_non_exhaustive()
}
}
impl SshTransportRes {
pub fn new(transport: impl SshTransport) -> Self {
Self { inner: Box::new(transport), generation: NEXT_GENERATION.fetch_add(1, Ordering::Relaxed) }
}
pub(crate) fn generation(&self) -> u64 {
self.generation
}
pub(crate) fn get_mut(&mut self) -> &mut dyn SshTransport {
self.inner.as_mut()
}
#[cfg_attr(not(feature = "sftp"), allow(dead_code))]
pub(crate) fn supports_sftp(&self) -> bool {
self.inner.supports_sftp()
}
}
#[derive(Clone)]
struct Script {
output: Vec<(SshStream, Vec<u8>)>,
result: Result<SshExit, BackendError>,
}
#[derive(Default)]
struct FakeState {
connects: Vec<(SshConnId, SshTarget)>,
live: HashSet<SshConnId>,
manual: bool,
reject_next: VecDeque<BackendError>,
scripts: HashMap<String, Script>,
commands: Vec<(SshConnId, RequestId, SshCommand)>,
running: HashSet<RequestId>,
#[cfg(feature = "sftp")]
sftp: Vec<(SshConnId, RequestId, SftpOp)>,
#[cfg(feature = "sftp")]
sftp_scripts: VecDeque<Result<SftpOutcome, BackendError>>,
cancelled: Vec<RequestId>,
closed: Vec<SshConnId>,
events: VecDeque<SshEvent>,
shutdowns: usize,
}
#[derive(Clone, Default)]
pub struct FakeSshTransport {
state: Arc<Mutex<FakeState>>,
}
impl fmt::Debug for FakeSshTransport {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let state = self.lock();
f.debug_struct("FakeSshTransport").field("connects", &state.connects.len()).field("commands", &state.commands.len()).finish_non_exhaustive()
}
}
impl FakeSshTransport {
pub const FINGERPRINT: &'static str = "SHA256:fake0fake0fake0fake0fake0fake0fake0fake0fa";
pub fn new() -> Self {
Self::default()
}
fn lock(&self) -> MutexGuard<'_, FakeState> {
self.state.lock().unwrap_or_else(PoisonError::into_inner)
}
fn event(&self, event: SshEvent) {
self.lock().events.push_back(event);
}
pub fn manual_connect(&self, manual: bool) -> &Self {
self.lock().manual = manual;
self
}
pub fn reject_next(&self, error: BackendError) -> &Self {
self.lock().reject_next.push_back(error);
self
}
pub fn on_command(&self, command: &str, output: &[(SshStream, &str)], result: Result<SshExit, BackendError>) -> &Self {
let output = output.iter().map(|(stream, text)| (*stream, text.as_bytes().to_vec())).collect();
self.lock().scripts.insert(command.to_string(), Script { output, result });
self
}
#[cfg(feature = "sftp")]
pub fn on_next_sftp(&self, result: Result<SftpOutcome, BackendError>) -> &Self {
self.lock().sftp_scripts.push_back(result);
self
}
pub fn accept(&self, conn: SshConnId) {
self.event(SshEvent::Connected { conn, fingerprint: Self::FINGERPRINT.to_string() });
}
pub fn drop_conn(&self, conn: SshConnId, error: BackendError) {
let mut state = self.lock();
state.live.remove(&conn);
state.events.push_back(SshEvent::Closed { conn, error: Some(error) });
}
pub fn output(&self, id: RequestId, stream: SshStream, data: &[u8]) {
self.event(SshEvent::Output { id, stream, data: data.to_vec() });
}
pub fn finish(&self, id: RequestId, result: Result<SshExit, BackendError>) {
let mut state = self.lock();
state.running.remove(&id);
state.events.push_back(SshEvent::Finished { id, result });
}
#[cfg(feature = "sftp")]
pub fn sftp_progress(&self, id: RequestId, done: u64, total: Option<u64>) {
self.event(SshEvent::Progress { id, done, total });
}
#[cfg(feature = "sftp")]
pub fn sftp_finish(&self, id: RequestId, result: Result<SftpOutcome, BackendError>) {
self.event(SshEvent::SftpFinished { id, result });
}
pub fn connects(&self) -> Vec<(SshConnId, SshTarget)> {
self.lock().connects.clone()
}
pub fn last_conn(&self) -> Option<SshConnId> {
self.lock().connects.last().map(|(conn, _)| *conn)
}
pub fn live_conns(&self) -> Vec<SshConnId> {
let mut conns: Vec<SshConnId> = self.lock().live.iter().copied().collect();
conns.sort_unstable();
conns
}
pub fn commands(&self) -> Vec<(SshConnId, RequestId, SshCommand)> {
self.lock().commands.clone()
}
pub fn running(&self) -> Vec<RequestId> {
let mut running: Vec<RequestId> = self.lock().running.iter().copied().collect();
running.sort_unstable();
running
}
#[cfg(feature = "sftp")]
pub fn sftp_ops(&self) -> Vec<(SshConnId, RequestId, SftpOp)> {
self.lock().sftp.clone()
}
pub fn cancelled(&self) -> Vec<RequestId> {
self.lock().cancelled.clone()
}
pub fn closed(&self) -> Vec<SshConnId> {
self.lock().closed.clone()
}
pub fn shutdown_count(&self) -> usize {
self.lock().shutdowns
}
}
impl SshTransport for FakeSshTransport {
fn connect(&mut self, conn: SshConnId, target: SshTarget) {
let mut state = self.lock();
state.connects.push((conn, target));
if let Some(error) = state.reject_next.pop_front() {
state.events.push_back(SshEvent::Closed { conn, error: Some(error) });
return;
}
state.live.insert(conn);
if !state.manual {
state.events.push_back(SshEvent::Connected { conn, fingerprint: Self::FINGERPRINT.to_string() });
}
}
fn run(&mut self, conn: SshConnId, id: RequestId, command: SshCommand) {
let mut state = self.lock();
let script = state.scripts.get(command.command()).cloned();
state.commands.push((conn, id, command));
if !state.live.contains(&conn) {
state.events.push_back(SshEvent::Finished { id, result: Err(BackendError::disconnected("not connected", Some(false))) });
return;
}
state.events.push_back(SshEvent::Started { id });
match script {
Some(script) => {
for (stream, data) in script.output {
state.events.push_back(SshEvent::Output { id, stream, data });
}
state.events.push_back(SshEvent::Finished { id, result: script.result });
}
None => {
state.running.insert(id);
}
}
}
#[cfg(feature = "sftp")]
fn sftp(&mut self, conn: SshConnId, id: RequestId, op: SftpOp) {
let mut state = self.lock();
state.sftp.push((conn, id, op));
if let Some(result) = state.sftp_scripts.pop_front() {
state.events.push_back(SshEvent::SftpFinished { id, result });
}
}
fn supports_sftp(&self) -> bool {
cfg!(feature = "sftp")
}
fn cancel(&mut self, id: RequestId) {
let mut state = self.lock();
state.running.remove(&id);
state.cancelled.push(id);
}
fn close(&mut self, conn: SshConnId) {
let mut state = self.lock();
state.live.remove(&conn);
state.closed.push(conn);
}
fn poll(&mut self) -> Vec<SshEvent> {
self.lock().events.drain(..).collect()
}
fn shutdown(&mut self) {
let mut state = self.lock();
state.shutdowns += 1;
state.live.clear();
state.running.clear();
}
}