#![cfg_attr(docsrs, feature(doc_cfg))]
#![warn(missing_docs)]
#![doc = include_str!(concat!(env!("CARGO_MANIFEST_DIR"), "/", env!("CARGO_PKG_README")))]
use std::collections::HashMap;
use std::convert::Infallible as Never;
use std::ops::Deref;
use std::path::Path;
use std::sync::Arc;
use russh::client::{connect, Config, Handle, Handler, Msg};
use russh::keys::{load_secret_key, ssh_key, PrivateKeyWithHashAlg};
use russh::{ChannelMsg, ChannelWriteHalf, CryptoVec};
use tokio::io::AsyncWrite;
use tokio::net::ToSocketAddrs;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
#[cfg(feature = "sftp")]
#[cfg_attr(docsrs, doc(cfg(feature = "sftp")))]
pub mod sftp;
use async_promise::Promise;
#[doc(no_inline)]
pub use russh::Error as SshError;
#[cfg(feature = "sftp")]
#[cfg_attr(docsrs, doc(cfg(feature = "sftp")))]
pub use russh_sftp;
use tracing::Instrument;
pub use {async_promise, russh, tokio};
mod read_stream;
pub use read_stream::ReadStream;
pub struct NoCheckHandler;
impl Handler for NoCheckHandler {
type Error = SshError;
async fn check_server_key(&mut self, _server_public_key: &ssh_key::PublicKey) -> Result<bool, Self::Error> {
Ok(true)
}
}
pub struct AsyncSession<H: Handler> {
session: Handle<H>,
}
impl<H: 'static + Handler> AsyncSession<H> {
pub async fn connect_unauthenticated(
config: Arc<Config>,
addrs: impl ToSocketAddrs,
handler: H,
) -> Result<Self, H::Error> {
let session = connect(config, addrs, handler).await?;
Ok(Self { session })
}
pub async fn open_channel(&self) -> Result<AsyncChannel, SshError> {
let russh_channel = self.session.channel_open_session().await?;
Ok(AsyncChannel::from(russh_channel))
}
}
impl AsyncSession<NoCheckHandler> {
pub async fn connect_publickey(
config: impl Into<Arc<Config>>,
addrs: impl ToSocketAddrs,
user: impl Into<String>,
key_path: impl AsRef<Path>,
) -> Result<Self, SshError> {
let key_pair = load_secret_key(key_path, None)?;
let mut session = connect(config.into(), addrs, NoCheckHandler).await?;
let auth_res = session
.authenticate_publickey(
user,
PrivateKeyWithHashAlg::new(Arc::new(key_pair), session.best_supported_rsa_hash().await?.flatten()),
)
.await?;
if auth_res.success() {
Ok(Self { session })
} else {
Err(SshError::NotAuthenticated)
}
}
}
impl<H: Handler> Deref for AsyncSession<H> {
type Target = Handle<H>;
fn deref(&self) -> &Self::Target {
&self.session
}
}
pub struct AsyncChannel {
write_half: ChannelWriteHalf<Msg>,
subscribe_send: mpsc::UnboundedSender<(Option<u32>, mpsc::UnboundedSender<CryptoVec>)>,
success_failure: Promise<bool>,
eof: Promise<()>,
exit_status: Promise<u32>,
closed: Promise<Never>,
_reader: JoinHandle<()>,
}
impl From<russh::Channel<Msg>> for AsyncChannel {
fn from(inner: russh::Channel<Msg>) -> Self {
let (mut read_half, write_half) = inner.split();
let (mut resolve_success_failure, success_failure) = async_promise::channel();
let (mut resolve_eof, eof) = async_promise::channel();
let (mut resolve_exit_status, exit_status) = async_promise::channel();
let (resolve_closed, closed) = async_promise::channel();
let (subscribe_send, mut subscribe_recv) = mpsc::unbounded_channel();
let reader = async move {
let _resolve_closed_drop = resolve_closed;
type Subscribers = HashMap<Option<u32>, mpsc::UnboundedSender<CryptoVec>>;
let mut subscribers = Some(Subscribers::new());
#[tracing::instrument(level = "INFO", skip_all, fields(?ext))]
fn receive_data(subscribers: &Option<Subscribers>, ext: Option<u32>, data: CryptoVec) {
if let Some(subscribers) = &subscribers {
if let Some(send) = subscribers.get(&ext) {
if let Err(e) = send.send(data) {
tracing::warn!("Failed to send data to subscriber: {e}");
} else {
tracing::debug!("Successfully sent data to subscriber.");
}
} else {
tracing::debug!("No subscriber for ext, dropping data.");
}
} else {
tracing::warn!("Unexpectedly received data from server after receiving EOF.");
}
}
loop {
tokio::select! {
biased;
Some((ext, send)) = subscribe_recv.recv() => {
if let Some(subscribers) = &mut subscribers {
subscribers.insert(ext, send);
} else {
tracing::debug!(ext, "Received stream subscriber after EOF, ignoring.");
}
},
opt_msg = read_half.wait() => {
let Some(msg) = opt_msg else {
break;
};
tracing::info_span!("Message", ?msg).in_scope(|| {
match msg {
ChannelMsg::Data { data } => receive_data(&subscribers, None, data),
ChannelMsg::ExtendedData { data, ext } => receive_data(&subscribers, Some(ext), data),
ChannelMsg::Success | ChannelMsg::Failure => {
tracing::debug!("Resolving success/failure.");
let is_success = matches!(msg, ChannelMsg::Success);
if resolve_success_failure.resolve(is_success).is_err() {
tracing::warn!("Success/failure already resolved, ignoring.");
}
}
ChannelMsg::Eof => {
tracing::debug!("Resolving EOF and dropping stream subscribers.");
if resolve_eof.resolve(()).is_err() {
tracing::warn!("EOF already resolved, ignoring.");
}
drop(std::mem::take(&mut subscribers));
}
ChannelMsg::ExitStatus { exit_status } => {
tracing::debug!(exit_status, "Resolving exit status.");
if resolve_exit_status.resolve(exit_status).is_err() {
tracing::warn!("Exit status already resolved, ignoring.");
}
}
_ => {
tracing::trace!("Ignoring message.");
}
}
});
},
}
}
tracing::debug!("Channel read half finished, reader exiting.");
};
let reader = tokio::task::spawn(reader.instrument(tracing::info_span!("Reader")));
Self {
write_half,
subscribe_send,
success_failure,
eof,
exit_status,
closed,
_reader: reader,
}
}
}
impl AsyncChannel {
pub fn read_stream(&self, ext: Option<u32>) -> ReadStream {
let (send, recv) = mpsc::unbounded_channel();
let _ = self.subscribe_send.send((ext, send));
ReadStream::from_recv(recv)
}
pub fn stdout(&self) -> ReadStream {
self.read_stream(None)
}
pub fn stderr(&self) -> ReadStream {
self.read_stream(Some(1))
}
pub fn write_stream(&self, ext: Option<u32>) -> impl use<'_> + 'static + AsyncWrite {
self.write_half.make_writer_ext(ext)
}
pub fn stdin(&self) -> impl use<'_> + 'static + AsyncWrite {
self.write_stream(None)
}
pub fn recv_success_failure(&self) -> &Promise<bool> {
&self.success_failure
}
pub fn recv_eof(&self) -> &Promise<()> {
&self.eof
}
pub fn recv_exit_status(&self) -> &Promise<u32> {
&self.exit_status
}
pub fn closed(&self) -> &Promise<Never> {
&self.closed
}
}
impl Deref for AsyncChannel {
type Target = ChannelWriteHalf<Msg>;
fn deref(&self) -> &Self::Target {
&self.write_half
}
}