use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use anyhow::{Result, bail};
use bytes::Bytes;
use russh::server::{Auth, Handler, Msg, Session};
use russh::{Channel, ChannelId};
use tokio::sync::mpsc;
use crate::state::{
CHANNEL_CAPACITY, DownMsg, PendingSession, PtySize, SessionQueue, SharedPtySize, UpMsg,
set_shared_pty_size,
};
struct ChannelState {
up_tx: mpsc::Sender<UpMsg>,
up_rx: Option<mpsc::Receiver<UpMsg>>,
down_rx: Option<mpsc::Receiver<DownMsg>>,
down_tx: Option<mpsc::Sender<DownMsg>>,
announced: bool,
}
fn passwords_equal(expected: &str, presented: &str) -> bool {
let (expected, presented) = (expected.as_bytes(), presented.as_bytes());
if expected.len() != presented.len() {
return false;
}
expected
.iter()
.zip(presented.iter())
.fold(0u8, |acc, (left, right)| acc | (left ^ right))
== 0
}
async fn drive_wire(
mut down_rx: mpsc::Receiver<DownMsg>,
handle: russh::server::Handle,
id: ChannelId,
) {
while let Some(message) = down_rx.recv().await {
let result = match message {
DownMsg::Data(bytes) => handle.data(id, bytes).await.map(|_| ()).map_err(|_| ()),
DownMsg::Eof => handle.eof(id).await,
};
if result.is_err() {
break;
}
}
let _ = handle.close(id).await;
}
pub struct EphemeralHandler {
expected_user: Arc<String>,
expected_pass: Arc<String>,
queue: Arc<SessionQueue>,
pty_size: SharedPtySize,
channels: HashMap<ChannelId, ChannelState>,
peer_addr: Option<SocketAddr>,
authed_user: Option<String>,
}
impl Handler for EphemeralHandler {
type Error = anyhow::Error;
async fn auth_password(&mut self, user: &str, password: &str) -> Result<Auth> {
if user == self.expected_user.as_str() && passwords_equal(&self.expected_pass, password) {
self.authed_user = Some(user.to_string());
Ok(Auth::Accept)
} else {
Ok(Auth::reject())
}
}
async fn channel_open_session(
&mut self,
channel: Channel<Msg>,
reply: russh::server::ChannelOpenHandle,
_session: &mut Session,
) -> Result<()> {
let (up_tx, up_rx) = mpsc::channel(CHANNEL_CAPACITY);
let (down_tx, down_rx) = mpsc::channel(CHANNEL_CAPACITY);
self.channels.insert(
channel.id(),
ChannelState {
up_tx,
up_rx: Some(up_rx),
down_rx: Some(down_rx),
down_tx: Some(down_tx),
announced: false,
},
);
reply.accept().await;
Ok(())
}
async fn shell_request(&mut self, channel: ChannelId, session: &mut Session) -> Result<()> {
self.announce(channel, None, session).await
}
async fn exec_request(
&mut self,
channel: ChannelId,
data: &[u8],
session: &mut Session,
) -> Result<()> {
let command = String::from_utf8_lossy(data).into_owned();
self.announce(channel, Some(command), session).await
}
async fn data(
&mut self,
channel: ChannelId,
data: &[u8],
_session: &mut Session,
) -> Result<()> {
let Some(state) = self.channels.get(&channel) else {
bail!("data for unknown channel");
};
state
.up_tx
.send(UpMsg::Data(Bytes::copy_from_slice(data)))
.await
.map_err(|_| anyhow::anyhow!("pump went away"))?;
Ok(())
}
async fn channel_eof(&mut self, channel: ChannelId, session: &mut Session) -> Result<()> {
let eof_delivered = match self.channels.get(&channel) {
Some(state) => state.up_tx.send(UpMsg::Eof).await.is_ok(),
None => false,
};
if !eof_delivered {
session
.handle()
.close(channel)
.await
.map_err(|_| anyhow::anyhow!("wire gone during EOF close"))?;
self.channels.remove(&channel);
}
Ok(())
}
async fn channel_close(&mut self, channel: ChannelId, _session: &mut Session) -> Result<()> {
self.channels.remove(&channel);
Ok(())
}
async fn pty_request(
&mut self,
channel: ChannelId,
_term: &str,
col_width: u32,
row_height: u32,
_pix_width: u32,
_pix_height: u32,
_modes: &[(russh::Pty, u32)],
session: &mut Session,
) -> Result<()> {
self.queue_size(row_height, col_width);
let _ = session.channel_success(channel);
Ok(())
}
async fn window_change_request(
&mut self,
_channel: ChannelId,
col_width: u32,
row_height: u32,
_pix_width: u32,
_pix_height: u32,
_session: &mut Session,
) -> Result<()> {
self.queue_size(row_height, col_width);
Ok(())
}
async fn env_request(
&mut self,
channel: ChannelId,
_variable_name: &str,
_variable_value: &str,
session: &mut Session,
) -> Result<()> {
let _ = session.channel_success(channel);
Ok(())
}
}
impl EphemeralHandler {
fn queue_size(&self, rows: u32, cols: u32) {
set_shared_pty_size(&self.pty_size, PtySize::new(rows, cols));
}
async fn announce(
&mut self,
channel: ChannelId,
exec_command: Option<String>,
session: &mut Session,
) -> Result<()> {
let Some(state) = self.channels.get_mut(&channel) else {
bail!("request for unknown channel");
};
if state.announced {
let _ = session.channel_success(channel);
return Ok(());
}
state.announced = true;
let (Some(up_rx), Some(down_rx), Some(down_tx)) = (
state.up_rx.take(),
state.down_rx.take(),
state.down_tx.take(),
) else {
bail!("channel already announced");
};
let handle = session.handle();
tokio::spawn(drive_wire(down_rx, handle, channel));
self.queue.push(PendingSession {
exec_command,
username: self.authed_user.clone(),
peer_addr: self.peer_addr,
pty_size: Arc::clone(&self.pty_size),
up_rx,
down_tx,
});
let _ = session.channel_success(channel);
Ok(())
}
}
pub struct ServerFactory {
expected_user: Arc<String>,
expected_pass: Arc<String>,
queue: Arc<SessionQueue>,
}
impl ServerFactory {
pub fn new(user: String, password: String, queue: Arc<SessionQueue>) -> Self {
Self {
expected_user: Arc::new(user),
expected_pass: Arc::new(password),
queue,
}
}
}
impl russh::server::Server for ServerFactory {
type Handler = EphemeralHandler;
fn new_client(&mut self, peer_addr: Option<SocketAddr>) -> EphemeralHandler {
EphemeralHandler {
expected_user: Arc::clone(&self.expected_user),
expected_pass: Arc::clone(&self.expected_pass),
queue: Arc::clone(&self.queue),
pty_size: crate::state::shared_pty_size(),
channels: HashMap::new(),
peer_addr,
authed_user: None,
}
}
}