use std::{net::SocketAddr, path::PathBuf, sync::Arc, time::Duration};
use base64::prelude::*;
use ed25519_dalek::VerifyingKey;
use futures::future;
use russh::keys::PrivateKeyWithHashAlg;
use russh::keys::ssh_key;
use russh::{
Channel, ChannelId, ChannelOpenFailure, ChannelStream, MethodKind, MethodSet, Preferred, Pty,
};
use serde::{Deserialize, Serialize};
#[cfg(test)]
use tokio::net::TcpListener;
use tokio::net::TcpStream;
use tokio::sync::oneshot;
use tracing::*;
use crate::error::{Result, TransportError};
use crate::transport::tls::certs::ServerConfigSource;
const DEFAULT_SOCK_ADDR: &str = "[::]:4422";
const DEFAULT_SSH_USER: &str = "ubuntu";
const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(30);
const INACTIVITY_TIMEOUT: Duration = Duration::from_secs(5 * 60);
const MAX_AUTH_ATTEMPTS: usize = 1;
pub type ServerChannelStream = ChannelStream<russh::server::Msg>;
pub type ClientChannelStream = ChannelStream<russh::client::Msg>;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct ServerConfig {
pub listen: SocketAddr,
pub connection_limit: Option<usize>,
pub identity_key: Option<String>,
pub private_ed25519_identity_key_file: Option<PathBuf>,
pub expected_username: Option<String>,
pub client_auth_key: Option<String>,
pub banner: Option<String>,
pub client_banner: Option<String>,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
listen: DEFAULT_SOCK_ADDR.parse().unwrap(),
connection_limit: Default::default(),
identity_key: Default::default(),
private_ed25519_identity_key_file: Default::default(),
expected_username: Default::default(),
client_auth_key: Default::default(),
banner: Default::default(),
client_banner: Default::default(),
}
}
}
impl ServerConfig {
fn get_crypto_source(&self) -> Result<ServerConfigSource> {
if let Some(ref base64_key) = self.identity_key {
ServerConfigSource::from_identity_base64(base64_key)
} else if let Some(ref key_path) = self.private_ed25519_identity_key_file {
ServerConfigSource::from_pkcs8_pem_file(key_path)
} else {
Err(TransportError::config_err("no crypto source provided"))
}
}
fn build_server_config(&self) -> Result<russh::server::Config> {
let source = self.get_crypto_source()?;
let keypair = ssh_key::private::Ed25519Keypair::from_seed(&source.identity_seed());
let host_key: russh::keys::PrivateKey = keypair.into();
let mut config = russh::server::Config {
methods: MethodSet::from(&[MethodKind::PublicKey][..]),
inactivity_timeout: Some(INACTIVITY_TIMEOUT),
keepalive_interval: Some(KEEPALIVE_INTERVAL),
max_auth_attempts: MAX_AUTH_ATTEMPTS,
auth_rejection_time: Duration::from_secs(3),
auth_rejection_time_initial: Some(Duration::from_secs(3)),
keys: vec![host_key],
preferred: Preferred {
key: std::borrow::Cow::Borrowed(&[ssh_key::Algorithm::Ed25519]),
..Preferred::default()
},
..Default::default()
};
if let Some(ref banner) = self.banner {
config.server_id = russh::SshId::Standard(std::borrow::Cow::Owned(banner.clone()));
}
Ok(config)
}
pub fn get_id_pubkey(&self) -> Result<String> {
let crypto_source = self.get_crypto_source()?;
let public_id = crypto_source.public_identity();
Ok(BASE64_STANDARD.encode(&public_id[..]))
}
pub fn client_auth_pubkey(&self) -> Result<VerifyingKey> {
let key = self.client_auth_key.as_deref().ok_or_else(|| {
TransportError::config_err("no client_auth_key configured for ssh transport")
})?;
let source = ServerConfigSource::from_identity_base64(key)?;
VerifyingKey::from_bytes(&source.public_identity())
.map_err(|e| TransportError::config_err(format!("bad client_auth_key: {e}")))
}
pub fn expected_username(&self) -> String {
self.expected_username
.clone()
.unwrap_or_else(|| DEFAULT_SSH_USER.to_string())
}
}
impl crate::transport::GenerateServerConfig for ServerConfig {
fn generate_config<R: rand::CryptoRng + ?Sized>(mut self, rng: &mut R) -> Self {
if self.identity_key.is_none() && self.private_ed25519_identity_key_file.is_none() {
self.identity_key = Some(ServerConfigSource::generate(rng).to_base64());
}
if self.client_auth_key.is_none() {
self.client_auth_key = Some(ServerConfigSource::generate(rng).to_base64());
}
self
}
}
impl crate::types::Sufficiency for ServerConfig {
fn is_sufficient(&self) -> bool {
(self.identity_key.is_some() || self.private_ed25519_identity_key_file.is_some())
&& self.client_auth_key.is_some()
}
}
impl crate::transport::ExternalizeKeyMaterial for ServerConfig {
fn externalize_keys(
mut self,
dir: &std::path::Path,
) -> Result<(Self, Vec<crate::transport::GeneratedKeyMaterial>)> {
let generated = crate::transport::externalize_identity(
&mut self.identity_key,
&mut self.private_ed25519_identity_key_file,
dir,
"ssh_ed25519_identity.pem",
)?;
Ok((self, generated))
}
}
pub fn create_listener(options: &ServerConfig) -> Result<Arc<russh::server::Config>> {
Ok(Arc::new(options.build_server_config()?))
}
struct ConnectionHandler {
channel_tx: Option<oneshot::Sender<Channel<russh::server::Msg>>>,
expected_username: String,
expected_client_pubkey: VerifyingKey,
peer_addr: SocketAddr,
}
impl ConnectionHandler {
fn deny_channel_request(
session: &mut russh::server::Session,
channel: ChannelId,
) -> std::result::Result<(), russh::Error> {
session.channel_failure(channel)
}
fn key_matches(&self, offered: &ssh_key::PublicKey) -> bool {
let Some(ed25519_key) = offered.key_data().ed25519() else {
return false;
};
ed25519_key.0 == self.expected_client_pubkey.to_bytes()
}
}
impl russh::server::Handler for ConnectionHandler {
type Error = russh::Error;
async fn auth_none(
&mut self,
user: &str,
) -> std::result::Result<russh::server::Auth, Self::Error> {
warn!(
peer = %self.peer_addr,
user,
"rejected ssh auth attempt: `none` method is disallowed"
);
Ok(russh::server::Auth::reject())
}
async fn auth_publickey(
&mut self,
user: &str,
public_key: &ssh_key::PublicKey,
) -> std::result::Result<russh::server::Auth, Self::Error> {
if user != self.expected_username {
warn!(
peer = %self.peer_addr,
user,
"rejected ssh publickey auth attempt: unexpected username"
);
return Err(russh::Error::Disconnect);
}
if !self.key_matches(public_key) {
warn!(
peer = %self.peer_addr,
user,
"rejected ssh publickey auth attempt: key does not match the configured identity"
);
return Err(russh::Error::Disconnect);
}
Ok(russh::server::Auth::Accept)
}
async fn auth_password(
&mut self,
user: &str,
_password: &str,
) -> std::result::Result<russh::server::Auth, Self::Error> {
warn!(
peer = %self.peer_addr,
user,
"rejected ssh auth attempt: `password` method is not supported"
);
Ok(russh::server::Auth::reject())
}
async fn auth_keyboard_interactive(
&mut self,
user: &str,
_submethods: &str,
_response: Option<russh::server::Response<'_>>,
) -> std::result::Result<russh::server::Auth, Self::Error> {
warn!(
peer = %self.peer_addr,
user,
"rejected ssh auth attempt: `keyboard-interactive` method is not supported"
);
Ok(russh::server::Auth::reject())
}
async fn auth_openssh_certificate(
&mut self,
user: &str,
_certificate: &ssh_key::Certificate,
) -> std::result::Result<russh::server::Auth, Self::Error> {
warn!(
peer = %self.peer_addr,
user,
"rejected ssh auth attempt: openssh certificate auth is not supported"
);
Ok(russh::server::Auth::reject())
}
async fn channel_open_session(
&mut self,
channel: Channel<russh::server::Msg>,
reply: russh::server::ChannelOpenHandle,
_session: &mut russh::server::Session,
) -> std::result::Result<(), Self::Error> {
reply.accept().await;
if let Some(tx) = self.channel_tx.take() {
let _ = tx.send(channel);
}
Ok(())
}
async fn channel_open_x11(
&mut self,
_channel: Channel<russh::server::Msg>,
_originator_address: &str,
_originator_port: u32,
reply: russh::server::ChannelOpenHandle,
_session: &mut russh::server::Session,
) -> std::result::Result<(), Self::Error> {
reply
.reject(ChannelOpenFailure::AdministrativelyProhibited)
.await;
Ok(())
}
async fn channel_open_direct_tcpip(
&mut self,
_channel: Channel<russh::server::Msg>,
_host_to_connect: &str,
_port_to_connect: u32,
_originator_address: &str,
_originator_port: u32,
reply: russh::server::ChannelOpenHandle,
_session: &mut russh::server::Session,
) -> std::result::Result<(), Self::Error> {
reply
.reject(ChannelOpenFailure::AdministrativelyProhibited)
.await;
Ok(())
}
async fn tcpip_forward(
&mut self,
_address: &str,
_port: &mut u32,
_session: &mut russh::server::Session,
) -> std::result::Result<bool, Self::Error> {
Ok(false)
}
async fn pty_request(
&mut self,
channel: ChannelId,
_term: &str,
_col_width: u32,
_row_height: u32,
_pix_width: u32,
_pix_height: u32,
_modes: &[(Pty, u32)],
session: &mut russh::server::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_request(session, channel)
}
async fn x11_request(
&mut self,
channel: ChannelId,
_single_connection: bool,
_x11_auth_protocol: &str,
_x11_auth_cookie: &str,
_x11_screen_number: u32,
session: &mut russh::server::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_request(session, channel)
}
async fn shell_request(
&mut self,
channel: ChannelId,
session: &mut russh::server::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_request(session, channel)
}
async fn exec_request(
&mut self,
channel: ChannelId,
_data: &[u8],
session: &mut russh::server::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_request(session, channel)
}
async fn subsystem_request(
&mut self,
channel: ChannelId,
_name: &str,
session: &mut russh::server::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_request(session, channel)
}
}
pub async fn accept(
config: Arc<russh::server::Config>,
expected_username: String,
expected_client_pubkey: VerifyingKey,
stream: TcpStream,
) -> Result<ServerChannelStream> {
let peer_addr = stream.peer_addr()?;
let (channel_tx, channel_rx) = oneshot::channel();
let handler = ConnectionHandler {
channel_tx: Some(channel_tx),
expected_username,
expected_client_pubkey,
peer_addr,
};
let running = russh::server::run_stream(config, stream, handler).await?;
tokio::spawn(async move {
if let Err(err) = running.await {
warn!("ssh session ended with error: {err}");
}
});
let channel = channel_rx
.await
.map_err(|_| TransportError::other("ssh session closed before a channel was opened"))?;
Ok(channel.into_stream())
}
pub use crate::types::ssh::ClientOptions;
struct InnerClientOptions {
pub addresses: Vec<SocketAddr>,
pub id_pubkey: VerifyingKey,
pub username: Option<String>,
pub client_banner: Option<String>,
pub client_auth_key: russh::keys::PrivateKey,
}
impl TryFrom<&ClientOptions> for InnerClientOptions {
type Error = TransportError;
fn try_from(value: &ClientOptions) -> Result<Self> {
let id_pubkey = Self::parse_base64_pubkey(&value.id_pubkey)?;
let client_auth_key = Self::parse_base64_auth_key(&value.client_auth_key)?;
Ok(Self {
addresses: value.addresses.clone(),
id_pubkey,
username: value.username.clone(),
client_banner: value.client_banner.clone(),
client_auth_key,
})
}
}
impl InnerClientOptions {
fn parse_base64_pubkey(key: impl AsRef<str>) -> Result<VerifyingKey> {
let mut pubkey_bytes = [0u8; 32];
BASE64_STANDARD
.decode_slice(key.as_ref(), &mut pubkey_bytes)
.map_err(|e| {
TransportError::config_err(format!(
"failed to decode SSH bridge public key as base64: {e}"
))
})?;
VerifyingKey::from_bytes(&pubkey_bytes)
.map_err(|e| TransportError::config_err(format!("bad SSH bridge public key: {e}")))
}
fn parse_base64_auth_key(key: impl AsRef<str>) -> Result<russh::keys::PrivateKey> {
let source = ServerConfigSource::from_identity_base64(key.as_ref())?;
let keypair = ssh_key::private::Ed25519Keypair::from_seed(&source.identity_seed());
Ok(keypair.into())
}
}
struct ClientHandler {
id_pubkey: VerifyingKey,
}
impl ClientHandler {
async fn deny_channel_open(reply: russh::client::ChannelOpenHandle) {
reply
.reject(ChannelOpenFailure::AdministrativelyProhibited)
.await;
}
}
impl russh::client::Handler for ClientHandler {
type Error = russh::Error;
async fn check_server_key(
&mut self,
server_public_key: &ssh_key::PublicKey,
) -> std::result::Result<bool, Self::Error> {
let Some(ed25519_key) = server_public_key.key_data().ed25519() else {
return Ok(false);
};
Ok(ed25519_key.0 == self.id_pubkey.to_bytes())
}
async fn should_accept_unknown_server_channel(
&mut self,
_id: ChannelId,
_channel_type: &str,
) -> bool {
false
}
async fn server_channel_open_session(
&mut self,
_channel: Channel<russh::client::Msg>,
reply: russh::client::ChannelOpenHandle,
_session: &mut russh::client::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_open(reply).await;
Ok(())
}
async fn server_channel_open_x11(
&mut self,
_channel: Channel<russh::client::Msg>,
_originator_address: &str,
_originator_port: u32,
reply: russh::client::ChannelOpenHandle,
_session: &mut russh::client::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_open(reply).await;
Ok(())
}
async fn server_channel_open_direct_tcpip(
&mut self,
_channel: Channel<russh::client::Msg>,
_host_to_connect: &str,
_port_to_connect: u32,
_originator_address: &str,
_originator_port: u32,
reply: russh::client::ChannelOpenHandle,
_session: &mut russh::client::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_open(reply).await;
Ok(())
}
async fn server_channel_open_direct_streamlocal(
&mut self,
_channel: Channel<russh::client::Msg>,
_socket_path: &str,
reply: russh::client::ChannelOpenHandle,
_session: &mut russh::client::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_open(reply).await;
Ok(())
}
async fn server_channel_open_forwarded_tcpip(
&mut self,
_channel: Channel<russh::client::Msg>,
_connected_address: &str,
_connected_port: u32,
_originator_address: &str,
_originator_port: u32,
reply: russh::client::ChannelOpenHandle,
_session: &mut russh::client::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_open(reply).await;
Ok(())
}
async fn server_channel_open_forwarded_streamlocal(
&mut self,
_channel: Channel<russh::client::Msg>,
_socket_path: &str,
reply: russh::client::ChannelOpenHandle,
_session: &mut russh::client::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_open(reply).await;
Ok(())
}
async fn server_channel_open_agent_forward(
&mut self,
_channel: Channel<russh::client::Msg>,
reply: russh::client::ChannelOpenHandle,
_session: &mut russh::client::Session,
) -> std::result::Result<(), Self::Error> {
Self::deny_channel_open(reply).await;
Ok(())
}
}
fn build_client_config(client_banner: Option<String>) -> russh::client::Config {
let mut config = russh::client::Config {
keepalive_interval: Some(KEEPALIVE_INTERVAL),
inactivity_timeout: Some(INACTIVITY_TIMEOUT),
preferred: Preferred {
key: std::borrow::Cow::Borrowed(&[ssh_key::Algorithm::Ed25519]),
..Preferred::default()
},
..Default::default()
};
if let Some(client_banner) = client_banner {
config.client_id = russh::SshId::Standard(std::borrow::Cow::Owned(client_banner));
}
config
}
pub async fn transport_conn(
options: &ClientOptions,
connect_timeout: Duration,
) -> Result<ClientChannelStream> {
info!("initializing from transport identity pubkey");
let inner_options = InnerClientOptions::try_from(options)?;
if inner_options.addresses.is_empty() {
return Err(TransportError::config_err(
"no ssh bridge address configured",
));
}
let client_config = Arc::new(build_client_config(inner_options.client_banner));
let handler = ClientHandler {
id_pubkey: inner_options.id_pubkey,
};
let connect = async {
let attempts = inner_options
.addresses
.iter()
.map(|&addr| Box::pin(TcpStream::connect(addr)));
let (stream, _losing_attempts) = future::select_ok(attempts).await.inspect_err(|e| {
warn!(
"failed to connect to any of {} ssh endpoint(s): {e}",
inner_options.addresses.len()
);
})?;
let mut handle = russh::client::connect_stream(client_config, stream, handler).await?;
let username = inner_options.username.unwrap_or(DEFAULT_SSH_USER.into());
let auth_key = PrivateKeyWithHashAlg::new(Arc::new(inner_options.client_auth_key), None);
let auth = handle.authenticate_publickey(&username, auth_key).await?;
if !auth.success() {
return Err(TransportError::other("ssh server rejected authentication"));
}
let channel = handle
.channel_open_session()
.await
.map_err(|e| TransportError::other(format!("failed to open ssh channel: {e}")))?;
Ok(channel.into_stream())
};
match tokio::time::timeout(connect_timeout, connect).await {
Ok(result) => result,
Err(_) => {
warn!("SSH bridge connection timed out after {connect_timeout:?}");
Err(TransportError::TimedOut(connect_timeout))
}
}
}
#[cfg(test)]
mod test;