diff --git a/src/main.rs b/src/main.rs index fdc74ad..dc14b45 100644 --- a/src/main.rs +++ b/src/main.rs @@ -5,7 +5,7 @@ use core::{ use std::{ borrow::Cow, ffi::OsString, - io::{BufReader, Read as _, Stdout, Write as _, stdout}, + io::{BufReader, IsTerminal as _, Read as _, Stdout, Write as _, stdout}, mem, path::PathBuf, sync::{Arc, Mutex}, @@ -139,44 +139,50 @@ async fn inner_main() -> Result<(), WrappedErr> { .canonicalize() .map_err(|e| WrappedErr(format!("Cannot canonicalize provided file: {e}").into()))?; - let (default_black, default_white) = if flags.terminal_colors { + if flags.terminal_colors && (flags.black_color.is_some() || flags.white_color.is_some()) { + return Err(WrappedErr( + "--terminal-colors cannot be combined with --black-color or --white-color".into() + )); + } + + let (black, white) = if flags.terminal_colors { query_terminal_colors() } else { - (MUPDF_BLACK, MUPDF_WHITE) + let black = flags + .black_color + .as_deref() + .map(|color| { + parse_color_to_i32(color).map_err(|e| { + WrappedErr( + format!( + "Couldn't parse black color {color:?}: {e} - is it formatted like a CSS color?" + ) + .into() + ) + }) + }) + .transpose()? + .unwrap_or(MUPDF_BLACK); + + let white = flags + .white_color + .as_deref() + .map(|color| { + parse_color_to_i32(color).map_err(|e| { + WrappedErr( + format!( + "Couldn't parse white color {color:?}: {e} - is it formatted like a CSS color?" + ) + .into() + ) + }) + }) + .transpose()? + .unwrap_or(MUPDF_WHITE); + + (black, white) }; - let black = flags - .black_color - .as_deref() - .map(|color| { - parse_color_to_i32(color).map_err(|e| { - WrappedErr( - format!( - "Couldn't parse black color {color:?}: {e} - is it formatted like a CSS color?" - ) - .into() - ) - }) - }) - .transpose()? - .unwrap_or(default_black); - - let white = flags - .white_color - .as_deref() - .map(|color| { - parse_color_to_i32(color).map_err(|e| { - WrappedErr( - format!( - "Couldn't parse white color {color:?}: {e} - is it formatted like a CSS color?" - ) - .into() - ) - }) - }) - .transpose()? - .unwrap_or(default_white); - // need to keep it around throughout the lifetime of the program, but don't rly need to use it. // Just need to make sure it doesn't get dropped yet. let maybe_logger = if std::env::var("RUST_LOG").is_ok() { @@ -574,42 +580,66 @@ fn parse_color_to_i32(cs: &str) -> Result } fn query_terminal_colors() -> (i32, i32) { + if !std::io::stdin().is_terminal() || !std::io::stdout().is_terminal() { + return (MUPDF_BLACK, MUPDF_WHITE); + } + let Ok(()) = enable_raw_mode() else { return (MUPDF_BLACK, MUPDF_WHITE); }; - let fg = query_osc_color(10); - let bg = query_osc_color(11); + struct RawModeGuard; + impl Drop for RawModeGuard { + fn drop(&mut self) { + let _ = disable_raw_mode(); + } + } + let _guard = RawModeGuard; - let _ = disable_raw_mode(); + let stdin = std::io::stdin(); + let mut handle = stdin.lock(); + + let fg = query_osc_color(10, &mut handle); + let bg = query_osc_color(11, &mut handle); + drop(handle); (fg.unwrap_or(MUPDF_BLACK), bg.unwrap_or(MUPDF_WHITE)) } -fn query_osc_color(osc: u8) -> Option { +fn query_osc_color(osc: u8, handle: &mut std::io::StdinLock<'_>) -> Option { print!("\x1b]{osc};?\x1b\\"); - std::io::stdout().flush().unwrap(); + std::io::stdout().flush().ok()?; - let stdin = std::io::stdin(); - let mut handle = stdin.lock(); - let mut buf = Vec::new(); + let mut buf = Vec::with_capacity(64); + let mut prev = None::; let mut byte = [0u8; 1]; loop { handle.read_exact(&mut byte).ok()?; - if byte[0] == b'\\' || byte[0] == 0x07 { + let b = byte[0]; + + if b == 0x07 || b == 0x9c { + break; + } + if prev == Some(0x1b) && b == b'\\' { + buf.pop(); break; } - buf.push(byte[0]); - } - drop(handle); - let input = String::from_utf8(buf).ok()?; - let rgb_str = input.split("rgb:").nth(1)?.trim_end_matches('\x1b'); + buf.push(b); + prev = Some(b); + } + + let input = core::str::from_utf8(&buf).ok()?; + let rgb_str = input.split("rgb:").nth(1)?; let mut parts = rgb_str.split('/'); let parse = |hex: &str| -> Option { let val = u16::from_str_radix(hex, 16).ok()?; - Some(if hex.len() <= 2 { val as u8 } else { (val >> 8) as u8 }) + Some(if hex.len() <= 2 { + val as u8 + } else { + (val >> 8) as u8 + }) }; let r = parse(parts.next()?)?;