use crate::char_buffer::CharBuffer;
use crate::errors::{ReplayError, ReplayResult};
use crate::session::Session;
use crossterm::terminal;
use portable_pty::{Child, CommandBuilder, NativePtySystem, PtySize, PtySystem};
use regex::Regex;
use std::io::{Read, Write};
use std::thread::{self, JoinHandle};
use std::time::Duration;
type Reader = Box<dyn Read + Send>;
type Writer = Box<dyn Write + Send>;
type ChildProc = Box<dyn Child + Send + Sync>;
pub fn run_internal<R: Read, W: Write + Send + 'static>(
user_input: R, user_output: W, record_user_input: bool, session_description: Option<String>, no_compression: bool, ) -> ReplayResult<()> {
terminal::enable_raw_mode()?;
let (pty_stdout, pty_stdin, child) = spawn_shell()?;
let output_reader = thread::spawn(move || read_from_pty(pty_stdout, user_output));
let exit_msg = handle_user_input(
user_input,
pty_stdin,
child,
record_user_input,
session_description,
no_compression,
)?;
terminal::disable_raw_mode()?;
join_output_thread(output_reader)?;
if record_user_input {
println!("{}", exit_msg);
}
Ok(())
}
fn spawn_shell() -> ReplayResult<(Reader, Writer, ChildProc)> {
let pty_system = NativePtySystem::default();
let pty_pair = pty_system.openpty(PtySize {
rows: 24,
cols: 80,
pixel_width: 0,
pixel_height: 0,
})?;
let bash_cmd = CommandBuilder::new("/bin/bash");
let bash_process = pty_pair.slave.spawn_command(bash_cmd)?;
drop(pty_pair.slave);
let pty_stdout = pty_pair.master.try_clone_reader()?; let pty_stdin = pty_pair.master.take_writer()?; Ok((pty_stdout, pty_stdin, bash_process))
}
fn handle_user_input<R: Read, W: Write>(
mut user_input: R,
mut pty_stdin: W,
mut child: ChildProc,
record_input: bool,
session_description: Option<String>,
no_compression: bool,
) -> ReplayResult<String> {
let mut buf = [0u8; 1]; let mut char_buffer = CharBuffer::new();
let exit_re = Regex::new(r"^\s*exit\s*$").unwrap();
let mut session: Option<Session> = if record_input {
Some(Session::new(session_description)?)
} else {
None
};
loop {
if child.try_wait()?.is_some() {
break;
}
let n = user_input.read(&mut buf)?;
if n == 0 {
break; } else if n != 1 {
unreachable!("Unexpected read size, should be 1 in terminal raw mode!");
}
let c = buf[0];
match c {
b'\x7F' => {
char_buffer.pop_char();
}
b'\x17' => {
char_buffer.pop_word();
}
b'\x03' => {
if let Some(sess) = session.as_mut() {
sess.remove_last_command();
}
char_buffer.clear();
}
b'\r' => {
char_buffer.push_char(b'\r');
if let Some(sess) = session.as_mut() {
sess.add_command(char_buffer.get_buf().to_vec());
}
if char_buffer.get_buf() == b"q\r" {
child.kill()?;
session = None; break;
}
if exit_re.is_match(std::str::from_utf8(char_buffer.get_buf())?) {
drop(pty_stdin);
break;
}
char_buffer.clear();
}
_ => {
char_buffer.push_char(c);
} }
pty_stdin.write_all(&buf)?;
pty_stdin.flush()?;
}
if let Some(sess) = session {
sess.save_session(!no_compression)?;
return Ok(String::from("Session saved"));
}
Ok(String::from("No session saved"))
}
fn read_from_pty<R: Read + Send, W: Write + Send>(
mut pty_output: R, mut user_output: W, ) -> ReplayResult<()> {
let mut buffer = [0u8; 1024];
loop {
let n = pty_output.read(&mut buffer)?;
if n == 0 {
break;
}
user_output.write_all(&buffer[..n])?;
user_output.flush()?;
}
Ok(())
}
fn join_output_thread(output_thread: JoinHandle<ReplayResult<()>>) -> ReplayResult<()> {
output_thread
.join()
.map_err(|err| ReplayError::ThreadPanic(format!("`user_output` with \n {:?}", err)))??;
Ok(())
}
pub struct RawModeReader {
data: Vec<u8>,
pos: usize,
delay: Duration,
}
impl Default for RawModeReader {
fn default() -> Self {
Self {
data: Vec::new(),
pos: 0,
delay: Duration::from_millis(10),
}
}
}
impl RawModeReader {
pub fn with_input(input: &[u8]) -> Self {
Self {
data: input.to_vec(),
pos: 0,
delay: Duration::from_millis(10),
}
}
pub fn with_input_and_delay(input: &[u8], delay: Duration) -> Self {
Self {
data: input.to_vec(),
pos: 0,
delay,
}
}
}
impl std::io::Read for RawModeReader {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
if self.pos >= self.data.len() {
return Ok(0);
}
std::thread::sleep(self.delay);
buf[0] = self.data[self.pos];
self.pos += 1;
Ok(1)
}
}
#[cfg(test)]
mod test {
use super::*;
use crate::paths::clear_replay_dir;
use serial_test::serial;
use std::io::sink;
fn run_and_get_commands(input: &[u8]) -> Vec<String> {
let reader = RawModeReader::with_input(input);
let _ = run_internal(reader, sink(), true, None, false);
Session::load_last_session()
.map(|sess| {
sess.iter_commands()
.map(|s| s.to_string())
.collect::<Vec<_>>()
})
.unwrap_or_default()
}
#[test]
#[serial]
fn record_commands_with_ctrl_c() {
clear_replay_dir().unwrap();
let cmds = run_and_get_commands(b"echo test_ctrl_c\rsleep 5\r\x03exit\r");
assert_eq!(
cmds,
vec!["echo test_ctrl_c\r", "exit\r"],
"Expected only echo and exit to be saved when Ctrl+C is used"
);
}
#[test]
#[serial]
fn record_commands_with_q_enter() {
clear_replay_dir().unwrap();
let cmds = run_and_get_commands(b"echo q\rq\r");
assert!(
cmds.is_empty(),
"No session should be saved when quitting with q+Enter"
);
}
#[test]
#[serial]
fn record_commands_with_ctrl_w() {
clear_replay_dir().unwrap();
let cmds = run_and_get_commands(b"echo 1 2\x17\rexit\r");
assert_eq!(
cmds,
vec!["echo 1 \r", "exit\r"],
"Ctrl+W should delete the last word before saving"
);
}
#[test]
#[serial]
fn record_commands_with_all_control_chars() {
clear_replay_dir().unwrap();
let cmds = run_and_get_commands(b"ls\recho\x7Fo test\x17test\rexit\r");
assert_eq!(
cmds,
vec!["ls\r", "echo test\r", "exit\r"],
"Combination of Backspace + Ctrl+W should still produce valid commands"
);
let session = Session::load_last_session().unwrap();
assert!(
session.description.is_none(),
"Session description should remain None by default"
);
}
#[test]
#[serial]
fn record_exit_command_only() {
clear_replay_dir().unwrap();
let cmds = run_and_get_commands(b"echo exit\r exit \r");
assert_eq!(
cmds,
vec!["echo exit\r", " exit \r"],
"Expected echo and exit commands to be saved"
);
}
}