use std::sync::Arc;
use russh::server::{ChannelOpenHandle, Msg};
use russh::{Channel, ChannelId, ChannelOpenFailure, ChannelWriteHalf};
use tokio::sync::mpsc;
use crate::host::{ExecOutput, MachineHost};
use crate::input::{Outbound, SessionInput};
use crate::session::SessionChannel;
pub struct Destination {
pub machine: String,
pub host: String,
pub port: u16,
}
pub fn open<H: MachineHost>(
host: Arc<H>,
destination: Destination,
channel: Channel<Msg>,
reply: ChannelOpenHandle,
outbound: Arc<Outbound>,
refused: mpsc::UnboundedSender<ChannelId>,
) -> SessionChannel {
let (_, writer) = channel.split();
let (input, taken) = SessionInput::new();
let task = tokio::spawn(async move {
let Destination {
machine,
host: peer,
port,
} = destination;
match host.connect_tcp(&machine, &peer, port, taken).await {
Ok(output) => {
outbound.send(reply.accept()).await;
forward(output, writer, &outbound).await;
}
Err(e) => {
tracing::info!(
machine,
host = peer,
port,
error = %format!("{e:#}"),
"ssh forward could not connect"
);
let _ = refused.send(writer.id());
reply.reject(ChannelOpenFailure::ConnectFailed).await;
}
}
});
SessionChannel::running(input, task.abort_handle())
}
async fn forward(mut output: impl ExecOutput, writer: ChannelWriteHalf<Msg>, outbound: &Outbound) {
let mut eof_sent = false;
loop {
match output.recv().await {
Some(Ok(frame)) => {
if !frame.data.is_empty()
&& outbound.send(writer.data_bytes(frame.data)).await.is_err()
{
return;
}
if frame.eof && !eof_sent {
eof_sent = true;
let _ = outbound.send(writer.eof()).await;
}
if frame.done {
break;
}
}
Some(Err(e)) => {
tracing::debug!(error = %e, "ssh forward broke");
break;
}
None => break,
}
}
if !eof_sent {
let _ = outbound.send(writer.eof()).await;
}
let _ = outbound.send(writer.close()).await;
}