use super::protocol;
use std::sync::{
OnceLock,
atomic::{AtomicBool, Ordering},
};
static SEND_POLICY_ALLOW: OnceLock<AtomicBool> = OnceLock::new();
fn send_policy_allow() -> &'static AtomicBool {
SEND_POLICY_ALLOW.get_or_init(|| AtomicBool::new(false))
}
pub fn set_send_policy_allow(allow: bool) {
send_policy_allow().store(allow, Ordering::Release);
}
pub fn send_policy_is_allow() -> bool {
send_policy_allow().load(Ordering::Acquire)
}
pub struct UserInputApproval {
pub approved: bool,
pub wrote_newline: bool,
}
pub fn approve_user_input(
code: &str,
reply: &tokio::sync::oneshot::Sender<protocol::IpcResponse>,
) -> UserInputApproval {
if send_policy_is_allow() {
return UserInputApproval {
approved: !reply.is_closed(),
wrote_newline: false,
};
}
if reply.is_closed() {
return UserInputApproval {
approved: false,
wrote_newline: false,
};
}
let prompt = approval_prompt(code);
use std::io::Write;
print!("{prompt}");
let _ = std::io::stdout().flush();
let finish = |approved: bool| {
if approved {
print!("\r\n");
} else {
print!("\r\n# [arf] IPC send declined.\r\n");
}
let _ = std::io::stdout().flush();
UserInputApproval {
approved,
wrote_newline: true,
}
};
let was_raw = crossterm::terminal::is_raw_mode_enabled().unwrap_or(false);
if !was_raw && crossterm::terminal::enable_raw_mode().is_err() {
return finish(false);
}
struct RawModeGuard {
was_raw: bool,
}
impl Drop for RawModeGuard {
fn drop(&mut self) {
if !self.was_raw {
let _ = crossterm::terminal::disable_raw_mode();
}
}
}
let _raw_mode_guard = RawModeGuard { was_raw };
loop {
if reply.is_closed() {
return finish(false);
}
if let Err(error) = crossterm::event::poll(std::time::Duration::from_millis(50)) {
log::debug!("IPC send approval input failed: {error}");
return finish(false);
}
arf_libr::process_r_events();
super::poll_ipc_requests();
if !crossterm::event::poll(std::time::Duration::from_millis(0)).unwrap_or(false) {
continue;
}
match crossterm::event::read() {
Ok(crossterm::event::Event::Key(key)) => {
use crossterm::event::{KeyCode, KeyModifiers};
let approved = matches!(key.code, KeyCode::Char('y' | 'Y'))
&& matches!(key.modifiers, KeyModifiers::NONE | KeyModifiers::SHIFT)
&& !reply.is_closed();
while crossterm::event::poll(std::time::Duration::from_millis(0)).unwrap_or(false) {
if crossterm::event::read().is_err() {
break;
}
}
return finish(approved);
}
Ok(_) => continue,
Err(error) => {
log::debug!("IPC send approval input failed: {error}");
return finish(false);
}
}
}
}
fn approval_prompt(code: &str) -> String {
use crossterm::style::Stylize;
let escaped = user_input_display(code);
format!(
"{}\r\n {}\r\n{}{}",
"# [arf] IPC send request:".dark_cyan(),
escaped.yellow(),
"# [arf] ".dark_cyan(),
"Press y to approve, any other key declines: "
.yellow()
.bold(),
)
}
pub(crate) fn user_input_display(code: &str) -> String {
code.chars()
.map(|character| {
if character.is_ascii_graphic() || character == ' ' {
character.to_string()
} else {
character.escape_default().to_string()
}
})
.collect()
}
pub fn reject_user_input_not_approved(reply: tokio::sync::oneshot::Sender<protocol::IpcResponse>) {
let _ = reply.send(protocol::IpcResponse::error(
protocol::INPUT_NOT_APPROVED,
"IPC send was not approved".to_string(),
));
}
#[cfg(test)]
mod approval_display_tests {
fn strip_sgr(input: &str) -> String {
let mut output = String::new();
let mut chars = input.chars();
while let Some(character) = chars.next() {
if character == '\x1b' && chars.next() == Some('[') {
for character in chars.by_ref() {
if character == 'm' {
break;
}
}
} else {
output.push(character);
}
}
output
}
#[test]
fn approval_prompt_styles_and_separates_heading_code_and_confirmation() {
let prompt = super::approval_prompt(
r#"system("SHOULD_BE_VISIBLE")
next"#,
);
let plain = format!("{}<END>", strip_sgr(&prompt).replace("\r\n", "\n"));
insta::assert_snapshot!(
plain,
@r###"
# [arf] IPC send request:
system("SHOULD_BE_VISIBLE")\nnext
# [arf] Press y to approve, any other key declines: <END>
"###
);
}
#[test]
fn display_includes_content_after_old_preview_limit() {
let source = format!(
"{}{}",
"0123456789".repeat(24),
r#"system("SHOULD_BE_VISIBLE")"#
);
let display = super::user_input_display(&source);
insta::assert_snapshot!(display, @r###"012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789system("SHOULD_BE_VISIBLE")"###);
}
#[test]
fn display_escapes_control_characters() {
let source = "line\n\tline\r\0\x1b[31m";
let display = super::user_input_display(source);
insta::assert_snapshot!(display, @r###"line\n\tline\r\u{0}\u{1b}[31m"###);
assert!(!display.contains('\n'));
assert!(!display.contains('\r'));
assert!(!display.contains('\0'));
assert!(!display.contains('\x1b'));
}
}