#![allow(dead_code)]
#![cfg(unix)]
use portable_pty::{Child, CommandBuilder, PtySize, native_pty_system};
use regex::Regex;
use std::io::{Read, Write};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::thread::{self, JoinHandle};
use std::time::{Duration, Instant};
const DEFAULT_TIMEOUT_MS: u64 = 15000;
const DEFAULT_ROWS: u16 = 24;
const DEFAULT_COLS: u16 = 80;
#[derive(Clone)]
pub struct ScreenSnapshot {
pub lines: Vec<String>,
pub cursor_row: u16,
pub cursor_col: u16,
}
struct SharedState {
output_buffer: String,
running: bool,
screen: ScreenSnapshot,
}
pub struct Terminal {
state: Arc<Mutex<SharedState>>,
pty_writer: Arc<Mutex<Box<dyn Write + Send>>>,
_reader_handle: JoinHandle<()>,
child: Box<dyn Child + Send + Sync>,
shutdown: Arc<AtomicBool>,
}
impl Terminal {
pub fn spawn() -> Result<Self, String> {
Self::spawn_with_args(&[])
}
pub fn spawn_with_args(args: &[&str]) -> Result<Self, String> {
Self::spawn_with_size(args, DEFAULT_ROWS, DEFAULT_COLS)
}
pub fn spawn_with_size(args: &[&str], rows: u16, cols: u16) -> Result<Self, String> {
assert!(rows > 0 && cols > 0, "PTY size must be non-zero");
let bin_path = env!("CARGO_BIN_EXE_arf");
let pty_system = native_pty_system();
let pair = pty_system
.openpty(PtySize {
rows,
cols,
pixel_width: 0,
pixel_height: 0,
})
.map_err(|e| format!("Failed to open PTY: {}", e))?;
let mut cmd = CommandBuilder::new(bin_path);
let has_history_dir = args.contains(&"--history-dir");
if !has_history_dir {
cmd.arg("--no-history");
}
for arg in args {
cmd.arg(*arg);
}
let child = pair
.slave
.spawn_command(cmd)
.map_err(|e| format!("Failed to spawn arf: {}", e))?;
let pty_writer = pair
.master
.take_writer()
.map_err(|e| format!("Failed to get PTY writer: {}", e))?;
let mut pty_reader = pair
.master
.try_clone_reader()
.map_err(|e| format!("Failed to get PTY reader: {}", e))?;
drop(pair.slave);
let state = Arc::new(Mutex::new(SharedState {
output_buffer: String::new(),
running: true,
screen: ScreenSnapshot {
lines: vec![String::new(); rows as usize],
cursor_row: 0,
cursor_col: 0,
},
}));
let pty_writer = Arc::new(Mutex::new(pty_writer));
let pty_writer_clone = Arc::clone(&pty_writer);
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_clone = Arc::clone(&shutdown);
let state_clone = Arc::clone(&state);
let reader_handle = thread::spawn(move || {
let (query_tx, query_rx) = std::sync::mpsc::channel::<()>();
struct CursorQueryDetector {
query_tx: std::sync::mpsc::Sender<()>,
}
impl vt100::Callbacks for CursorQueryDetector {
fn unhandled_csi(
&mut self,
_screen: &mut vt100::Screen,
_prefix: Option<u8>,
_intermediate: Option<u8>,
params: &[&[u16]],
c: char,
) {
if c == 'n' {
let is_dsr = params.is_empty()
|| (params.len() == 1 && params[0].len() == 1 && params[0][0] == 6);
if is_dsr {
let _ = self.query_tx.send(());
}
}
}
}
let callbacks = CursorQueryDetector { query_tx };
let mut parser = vt100::Parser::new_with_callbacks(rows, cols, 0, callbacks);
let mut buf = [0u8; 4096];
loop {
if shutdown_clone.load(Ordering::Relaxed) {
break;
}
match pty_reader.read(&mut buf) {
Ok(0) => {
if let Ok(mut state) = state_clone.lock() {
state.running = false;
}
break;
}
Ok(n) => {
let data = &buf[..n];
parser.process(data);
if let Ok(mut state) = state_clone.lock() {
if let Ok(s) = std::str::from_utf8(data) {
state.output_buffer.push_str(s);
}
let screen = parser.screen();
let (cursor_row, cursor_col) = screen.cursor_position();
state.screen.cursor_row = cursor_row;
state.screen.cursor_col = cursor_col;
for row in 0..rows {
let row_content = screen.contents_between(row, 0, row, cols - 1);
state.screen.lines[row as usize] = row_content;
}
}
while query_rx.try_recv().is_ok() {
let (row, col) = parser.screen().cursor_position();
let response = format!("\x1b[{};{}R", row + 1, col + 1);
if let Ok(mut writer) = pty_writer_clone.lock() {
let _ = writer.write_all(response.as_bytes());
let _ = writer.flush();
}
}
}
Err(e) => {
if e.kind() != std::io::ErrorKind::WouldBlock
&& e.kind() != std::io::ErrorKind::Interrupted
{
if let Ok(mut state) = state_clone.lock() {
state.running = false;
}
break;
}
}
}
}
});
Ok(Terminal {
state,
pty_writer,
_reader_handle: reader_handle,
child,
shutdown,
})
}
pub fn expect(&mut self, pattern: &str) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
{
let state = self.state.lock().map_err(|e| e.to_string())?;
if !state.running && !state.output_buffer.contains(pattern) {
return Err(format!(
"Process exited before finding pattern '{}'. Output:\n{}",
pattern, state.output_buffer
));
}
if state.output_buffer.contains(pattern) {
return Ok(());
}
}
thread::sleep(Duration::from_millis(50));
}
let output = self
.state
.lock()
.map(|s| s.output_buffer.clone())
.unwrap_or_default();
Err(format!(
"Timeout waiting for pattern '{}'. Current output:\n{}",
pattern, output
))
}
#[allow(dead_code)]
pub fn expect_regex(&mut self, pattern: &str) -> Result<(), String> {
let re = Regex::new(pattern).map_err(|e| format!("Invalid regex: {}", e))?;
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
{
let state = self.state.lock().map_err(|e| e.to_string())?;
if !state.running && !re.is_match(&state.output_buffer) {
return Err(format!(
"Process exited before matching regex '{}'. Output:\n{}",
pattern, state.output_buffer
));
}
if re.is_match(&state.output_buffer) {
return Ok(());
}
}
thread::sleep(Duration::from_millis(50));
}
let output = self
.state
.lock()
.map(|s| s.output_buffer.clone())
.unwrap_or_default();
Err(format!(
"Timeout waiting for regex '{}'. Current output:\n{}",
pattern, output
))
}
pub fn wait_for_prompt(&mut self) -> Result<(), String> {
self.expect("> ")
}
pub fn clear_and_expect(&mut self, pattern: &str) -> Result<(), String> {
{
let mut state = self.state.lock().map_err(|e| e.to_string())?;
state.output_buffer.clear();
}
self.expect(pattern)
}
pub fn clear_buffer(&mut self) -> Result<(), String> {
let mut state = self.state.lock().map_err(|e| e.to_string())?;
state.output_buffer.clear();
Ok(())
}
pub fn get_output(&self) -> Result<String, String> {
let state = self.state.lock().map_err(|e| e.to_string())?;
Ok(state.output_buffer.clone())
}
pub fn send_line(&mut self, text: &str) -> Result<(), String> {
let data = format!("{}\n", text);
let mut writer = self.pty_writer.lock().map_err(|e| e.to_string())?;
writer
.write_all(data.as_bytes())
.map_err(|e| format!("Failed to send line: {}", e))?;
writer
.flush()
.map_err(|e| format!("Failed to flush: {}", e))
}
pub fn send(&mut self, text: &str) -> Result<(), String> {
let mut writer = self.pty_writer.lock().map_err(|e| e.to_string())?;
writer
.write_all(text.as_bytes())
.map_err(|e| format!("Failed to send: {}", e))?;
writer
.flush()
.map_err(|e| format!("Failed to flush: {}", e))
}
pub fn send_interrupt(&mut self) -> Result<(), String> {
self.send("\x03")
}
pub fn send_eof(&mut self) -> Result<(), String> {
self.send("\x04")
}
pub fn quit(&mut self) -> Result<(), String> {
let _ = self.send_line("q()");
thread::sleep(Duration::from_millis(500));
{
let state = self.state.lock().map_err(|e| e.to_string())?;
if state.running {
drop(state);
let _ = self.send_eof();
}
}
self.shutdown.store(true, Ordering::Relaxed);
let _ = self.child.kill();
Ok(())
}
pub fn screen(&self) -> Result<ScreenSnapshot, String> {
let state = self.state.lock().map_err(|e| e.to_string())?;
Ok(state.screen.clone())
}
pub fn line(&self, row: usize) -> ScreenLine {
ScreenLine {
state: Arc::clone(&self.state),
line_getter: LineGetter::Absolute(row),
}
}
pub fn current_line(&self) -> ScreenLine {
ScreenLine {
state: Arc::clone(&self.state),
line_getter: LineGetter::CurrentLine,
}
}
pub fn previous_line(&self, n: usize) -> ScreenLine {
ScreenLine {
state: Arc::clone(&self.state),
line_getter: LineGetter::PreviousLine(n),
}
}
pub fn process_id(&self) -> Option<u32> {
self.child.process_id()
}
pub fn cursor_position(&self) -> Result<(u16, u16), String> {
let state = self.state.lock().map_err(|e| e.to_string())?;
Ok((state.screen.cursor_row, state.screen.cursor_col))
}
pub fn assert_cursor(&self, expected_row: u16, expected_col: u16) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
let (row, col) = self.cursor_position()?;
if row == expected_row && col == expected_col {
return Ok(());
}
thread::sleep(Duration::from_millis(50));
}
let (row, col) = self.cursor_position()?;
Err(format!(
"Cursor position mismatch: expected ({}, {}), got ({}, {})",
expected_row, expected_col, row, col
))
}
#[allow(dead_code)]
pub fn dump_screen(&self) -> Result<(), String> {
let state = self.state.lock().map_err(|e| e.to_string())?;
eprintln!("=== Screen Dump ===");
eprintln!(
"Cursor: ({}, {})",
state.screen.cursor_row, state.screen.cursor_col
);
for (i, line) in state.screen.lines.iter().enumerate() {
let trimmed = line.trim_end();
if !trimmed.is_empty() || i == state.screen.cursor_row as usize {
let marker = if i == state.screen.cursor_row as usize {
">"
} else {
" "
};
eprintln!("{} {:2}: {:?}", marker, i, trimmed);
}
}
eprintln!("===================");
Ok(())
}
}
#[allow(dead_code)]
enum LineGetter {
Absolute(usize),
CurrentLine,
PreviousLine(usize),
}
pub struct ScreenLine {
state: Arc<Mutex<SharedState>>,
line_getter: LineGetter,
}
impl ScreenLine {
fn get_line(&self) -> Result<String, String> {
let state = self.state.lock().map_err(|e| e.to_string())?;
let row = match self.line_getter {
LineGetter::Absolute(r) => r,
LineGetter::CurrentLine => state.screen.cursor_row as usize,
LineGetter::PreviousLine(n) => (state.screen.cursor_row as usize).saturating_sub(n),
};
Ok(state.screen.lines.get(row).cloned().unwrap_or_default())
}
#[allow(dead_code)]
pub fn assert_startswith(&self, prefix: &str) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
let line = self.get_line()?;
if line.starts_with(prefix) {
return Ok(());
}
thread::sleep(Duration::from_millis(50));
}
let line = self.get_line()?;
Err(format!(
"Line does not start with '{}': got '{}'",
prefix, line
))
}
#[allow(dead_code)]
pub fn assert_endswith(&self, suffix: &str) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
let line = self.get_line()?;
if line.trim_end().ends_with(suffix) {
return Ok(());
}
thread::sleep(Duration::from_millis(50));
}
let line = self.get_line()?;
Err(format!(
"Line does not end with '{}': got '{}'",
suffix, line
))
}
pub fn assert_contains(&self, substring: &str) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
let line = self.get_line()?;
if line.contains(substring) {
return Ok(());
}
thread::sleep(Duration::from_millis(50));
}
let line = self.get_line()?;
Err(format!(
"Line does not contain '{}': got '{}'",
substring, line
))
}
#[allow(dead_code)]
pub fn assert_equal(&self, expected: &str) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
let line = self.get_line()?;
if line.trim() == expected.trim() {
return Ok(());
}
thread::sleep(Duration::from_millis(50));
}
let line = self.get_line()?;
Err(format!(
"Line does not equal '{}': got '{}'",
expected, line
))
}
#[allow(dead_code)]
pub fn trim(&self) -> TrimmedScreenLine<'_> {
TrimmedScreenLine { inner: self }
}
}
pub struct TrimmedScreenLine<'a> {
inner: &'a ScreenLine,
}
impl TrimmedScreenLine<'_> {
#[allow(dead_code)]
pub fn assert_startswith(&self, prefix: &str) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
let line = self.inner.get_line()?;
if line.trim().starts_with(prefix) {
return Ok(());
}
thread::sleep(Duration::from_millis(50));
}
let line = self.inner.get_line()?;
Err(format!(
"Trimmed line does not start with '{}': got '{}'",
prefix,
line.trim()
))
}
#[allow(dead_code)]
pub fn assert_endswith(&self, suffix: &str) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
let line = self.inner.get_line()?;
if line.trim().ends_with(suffix) {
return Ok(());
}
thread::sleep(Duration::from_millis(50));
}
let line = self.inner.get_line()?;
Err(format!(
"Trimmed line does not end with '{}': got '{}'",
suffix,
line.trim()
))
}
#[allow(dead_code)]
pub fn assert_equal(&self, expected: &str) -> Result<(), String> {
let timeout = Duration::from_millis(DEFAULT_TIMEOUT_MS);
let start = Instant::now();
while start.elapsed() < timeout {
let line = self.inner.get_line()?;
if line.trim() == expected {
return Ok(());
}
thread::sleep(Duration::from_millis(50));
}
let line = self.inner.get_line()?;
Err(format!(
"Trimmed line does not equal '{}': got '{}'",
expected,
line.trim()
))
}
}
impl Drop for Terminal {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::Relaxed);
let _ = self.child.kill();
}
}
use std::process::Command;
pub fn has_air_cli() -> bool {
Command::new("air")
.arg("--version")
.output()
.map(|o| o.status.success())
.unwrap_or(false)
}
pub fn has_dplyr() -> bool {
Command::new("Rscript")
.args(["-e", "library(dplyr)"])
.output()
.map(|o| o.status.success())
.unwrap_or(false)
}
pub fn has_askpass() -> bool {
Command::new("Rscript")
.args(["-e", "library(askpass)"])
.output()
.map(|o| o.status.success())
.unwrap_or(false)
}