use std::collections::HashMap;
use std::sync::Arc;
use arcbox_connect::v1::{MachineExecOutput, MachineExecRequest, TerminalSize};
use russh::ChannelWriteHalf;
use russh::server::{Handle, Msg};
use tokio::task::AbortHandle;
use crate::host::ExecOutput;
use crate::input::{Outbound, SessionInput};
use crate::signal;
use crate::target::Target;
pub const STDERR: u32 = 1;
pub const EXIT_FAILURE: u32 = 255;
pub enum Program {
Shell,
Command(String),
}
struct Pty {
term: String,
cols: u32,
rows: u32,
}
pub struct SessionChannel {
writer: Option<ChannelWriteHalf<Msg>>,
pty: Option<Pty>,
env: HashMap<String, String>,
input: Option<SessionInput>,
output_task: Option<AbortHandle>,
}
impl SessionChannel {
pub fn new(writer: ChannelWriteHalf<Msg>) -> Self {
Self {
writer: Some(writer),
pty: None,
env: HashMap::new(),
input: None,
output_task: None,
}
}
pub fn running(input: SessionInput, output_task: AbortHandle) -> Self {
Self {
writer: None,
pty: None,
env: HashMap::new(),
input: Some(input),
output_task: Some(output_task),
}
}
pub fn started(&self) -> bool {
self.writer.is_none()
}
pub fn request_pty(&mut self, term: &str, cols: u32, rows: u32) -> bool {
if self.started() {
return false;
}
self.pty = Some(Pty {
term: term.to_owned(),
cols,
rows,
});
true
}
pub fn request_env(&mut self, name: &str, value: &str) -> bool {
let valid = !name.is_empty() && !name.contains(['=', '\0']) && !value.contains('\0');
if self.started() || !valid {
return false;
}
self.env.insert(name.to_owned(), value.to_owned());
true
}
pub fn exec_request(
&self,
target: &Target,
program: Program,
ssh_env: &[(String, String)],
) -> MachineExecRequest {
let mut env = self.env.clone();
env.extend(ssh_env.iter().cloned());
if let Some(pty) = &self.pty {
env.insert("TERM".to_owned(), pty.term.clone());
}
let tty_size = self
.pty
.as_ref()
.filter(|pty| pty.cols > 0 && pty.rows > 0)
.map(|pty| TerminalSize {
width: pty.cols,
height: pty.rows,
..Default::default()
});
MachineExecRequest {
id: target.machine.clone(),
cmd: match program {
Program::Shell => Vec::new(),
Program::Command(line) => vec![line],
},
user: target.user.clone().unwrap_or_default(),
env: env.into_iter().collect(),
tty: self.pty.is_some(),
tty_size: tty_size.into(),
attach_stdin: true,
login: true,
..Default::default()
}
}
pub fn start<O: ExecOutput>(
&mut self,
started: anyhow::Result<(O, SessionInput)>,
handle: Handle,
outbound: Arc<Outbound>,
) {
let Some(writer) = self.writer.take() else {
return;
};
let newline = self.newline();
let task = match started {
Ok((output, input)) => {
self.input = Some(input);
tokio::spawn(forward_output(output, writer, handle, outbound, newline))
}
Err(e) => {
let exit = Exit::Failed(format!("{e:#}"));
tokio::spawn(finish(writer, handle, outbound, exit, newline, false))
}
};
self.output_task = Some(task.abort_handle());
}
pub const fn input(&self) -> Option<&SessionInput> {
self.input.as_ref()
}
pub const fn newline(&self) -> &'static str {
if self.pty.is_some() { "\r\n" } else { "\n" }
}
}
impl Drop for SessionChannel {
fn drop(&mut self) {
if let Some(task) = &self.output_task {
task.abort();
}
}
}
enum Exit {
Status(u32),
Signal(String),
Failed(String),
}
impl From<&MachineExecOutput> for Exit {
fn from(last: &MachineExecOutput) -> Self {
if last.exit_signal.is_empty() {
Self::Status(u32::try_from(last.exit_code).unwrap_or(EXIT_FAILURE))
} else {
Self::Signal(last.exit_signal.clone())
}
}
}
async fn forward_output(
mut output: impl ExecOutput,
writer: ChannelWriteHalf<Msg>,
handle: Handle,
outbound: Arc<Outbound>,
newline: &'static str,
) {
let mut eof_sent = false;
let exit = loop {
match output.recv().await {
Some(Ok(mut frame)) => {
if !frame.data.is_empty() {
let data = std::mem::take(&mut frame.data);
let sent = if frame.stream == "stderr" {
outbound
.send(writer.extended_data_bytes(STDERR, data))
.await
} else {
outbound.send(writer.data_bytes(data)).await
};
if sent.is_err() {
return;
}
}
if frame.eof && !eof_sent {
eof_sent = true;
let _ = outbound.send(writer.eof()).await;
}
if frame.done {
break Exit::from(&frame);
}
}
Some(Err(e)) => break Exit::Failed(e.to_string()),
None => break Exit::Failed("the machine session ended without an exit status".into()),
}
};
finish(writer, handle, outbound, exit, newline, eof_sent).await;
}
async fn finish(
writer: ChannelWriteHalf<Msg>,
handle: Handle,
outbound: Arc<Outbound>,
exit: Exit,
newline: &str,
eof_sent: bool,
) {
match exit {
Exit::Status(code) => {
let _ = outbound.send(writer.exit_status(code)).await;
}
Exit::Signal(name) => {
let sig = signal::from_name(&name);
let _ = outbound
.send(handle.exit_signal_request(
writer.id(),
sig,
false,
String::new(),
String::new(),
))
.await;
}
Exit::Failed(message) => {
if eof_sent {
tracing::info!(%message, "ssh session broke after its output ended");
} else {
let message = format!("arcbox: {message}{newline}");
let _ = outbound
.send(writer.extended_data_bytes(STDERR, message))
.await;
}
let _ = outbound.send(writer.exit_status(EXIT_FAILURE)).await;
}
}
if !eof_sent {
let _ = outbound.send(writer.eof()).await;
}
let _ = outbound.send(writer.close()).await;
}