use super::*;
use std::io;
const INPUT_POLL_IDLE_TIMEOUT: Duration = Duration::from_millis(100);
const INPUT_POLL_DRAIN_TIMEOUT: Duration = Duration::ZERO;
pub(crate) struct TerminalInputBridge {
pub(crate) receiver: Receiver<crossterm::event::Event>,
shutdown: Arc<AtomicBool>,
handle: Option<JoinHandle<()>>,
}
trait TerminalEventReader: Send + 'static {
fn poll(&mut self, timeout: Duration) -> io::Result<bool>;
fn read(&mut self) -> io::Result<crossterm::event::Event>;
}
struct CrosstermEventReader;
impl TerminalEventReader for CrosstermEventReader {
fn poll(&mut self, timeout: Duration) -> io::Result<bool> {
crossterm::event::poll(timeout)
}
fn read(&mut self) -> io::Result<crossterm::event::Event> {
crossterm::event::read()
}
}
impl TerminalInputBridge {
pub(crate) fn spawn() -> Self {
Self::spawn_with_reader(CrosstermEventReader)
}
fn spawn_with_reader<R>(reader: R) -> Self
where
R: TerminalEventReader,
{
let (sender, receiver) = bounded::<crossterm::event::Event>(1024);
let shutdown = Arc::new(AtomicBool::new(false));
let thread_shutdown = Arc::clone(&shutdown);
let handle = thread::spawn(move || {
run_input_reader(reader, sender, thread_shutdown);
});
Self {
receiver,
shutdown,
handle: Some(handle),
}
}
}
fn run_input_reader<R>(
mut reader: R,
sender: crossbeam_channel::Sender<crossterm::event::Event>,
shutdown: Arc<AtomicBool>,
) where
R: TerminalEventReader,
{
let mut poll_timeout = INPUT_POLL_IDLE_TIMEOUT;
while !shutdown.load(Ordering::SeqCst) {
match reader.poll(poll_timeout) {
Ok(true) => {
let event = match reader.read() {
Ok(event) => event,
Err(_) => break,
};
if shutdown.load(Ordering::SeqCst) {
break;
}
match sender.try_send(event) {
Ok(()) | Err(crossbeam_channel::TrySendError::Full(_)) => {}
Err(crossbeam_channel::TrySendError::Disconnected(_)) => break,
}
poll_timeout = INPUT_POLL_DRAIN_TIMEOUT;
}
Ok(false) => {
poll_timeout = INPUT_POLL_IDLE_TIMEOUT;
}
Err(_) => break,
}
}
}
impl Drop for TerminalInputBridge {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::SeqCst);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crossterm::event::{Event, KeyCode, KeyEvent, KeyEventKind, KeyModifiers};
use std::sync::Mutex;
fn key_event(ch: char) -> Event {
Event::Key(KeyEvent::new_with_kind(
KeyCode::Char(ch),
KeyModifiers::NONE,
KeyEventKind::Press,
))
}
struct ScriptedReader {
steps: VecDeque<io::Result<Option<Event>>>,
observed_timeouts: Arc<Mutex<Vec<Duration>>>,
}
impl ScriptedReader {
fn new(
steps: impl Into<VecDeque<io::Result<Option<Event>>>>,
observed_timeouts: Arc<Mutex<Vec<Duration>>>,
) -> Self {
Self {
steps: steps.into(),
observed_timeouts,
}
}
}
impl TerminalEventReader for ScriptedReader {
fn poll(&mut self, timeout: Duration) -> io::Result<bool> {
self.observed_timeouts.lock().unwrap().push(timeout);
match self.steps.front() {
Some(Ok(Some(_))) => Ok(true),
Some(Ok(None)) => {
self.steps.pop_front();
Ok(false)
}
Some(Err(_)) | None => Err(io::Error::other("reader stopped")),
}
}
fn read(&mut self) -> io::Result<Event> {
match self.steps.pop_front() {
Some(Ok(Some(event))) => Ok(event),
Some(Ok(None)) => Err(io::Error::other("no event ready")),
Some(Err(error)) => Err(error),
None => Err(io::Error::other("reader stopped")),
}
}
}
struct TimedPollReader {
first_poll_started: Arc<AtomicBool>,
poll_count: Arc<AtomicU64>,
}
impl TerminalEventReader for TimedPollReader {
fn poll(&mut self, _timeout: Duration) -> io::Result<bool> {
self.first_poll_started.store(true, Ordering::SeqCst);
self.poll_count.fetch_add(1, Ordering::SeqCst);
thread::sleep(Duration::from_millis(10));
Ok(false)
}
fn read(&mut self) -> io::Result<Event> {
Err(io::Error::other("read should not be called"))
}
}
struct ShutdownOnReadReader {
shutdown: Arc<AtomicBool>,
}
impl TerminalEventReader for ShutdownOnReadReader {
fn poll(&mut self, _timeout: Duration) -> io::Result<bool> {
Ok(true)
}
fn read(&mut self) -> io::Result<Event> {
self.shutdown.store(true, Ordering::SeqCst);
Ok(key_event('q'))
}
}
#[test]
fn drop_sets_shutdown_and_joins_reader_thread() {
let first_poll_started = Arc::new(AtomicBool::new(false));
let poll_count = Arc::new(AtomicU64::new(0));
let bridge = TerminalInputBridge::spawn_with_reader(TimedPollReader {
first_poll_started: Arc::clone(&first_poll_started),
poll_count: Arc::clone(&poll_count),
});
while !first_poll_started.load(Ordering::SeqCst) {
thread::sleep(Duration::from_millis(1));
}
let before_drop = poll_count.load(Ordering::SeqCst);
let started = Instant::now();
drop(bridge);
assert!(started.elapsed() < Duration::from_secs(1));
assert!(poll_count.load(Ordering::SeqCst) <= before_drop + 1);
}
#[test]
fn shutdown_after_read_drops_event_before_forwarding() {
let shutdown = Arc::new(AtomicBool::new(false));
let (sender, receiver) = bounded(1);
run_input_reader(
ShutdownOnReadReader {
shutdown: Arc::clone(&shutdown),
},
sender,
shutdown,
);
assert!(receiver.try_recv().is_err());
}
#[test]
fn forwards_events_and_drops_when_channel_is_full() {
let observed_timeouts = Arc::new(Mutex::new(Vec::new()));
let mut steps = VecDeque::new();
for _ in 0..1030 {
steps.push_back(Ok(Some(key_event('x'))));
}
steps.push_back(Err(io::Error::other("done")));
let bridge = TerminalInputBridge::spawn_with_reader(ScriptedReader::new(
steps,
Arc::clone(&observed_timeouts),
));
if let Some(handle) = bridge.handle.as_ref() {
while !handle.is_finished() {
thread::sleep(Duration::from_millis(1));
}
}
let received: Vec<_> = bridge.receiver.try_iter().collect();
assert_eq!(received.len(), 1024);
assert!(received.iter().all(|event| *event == key_event('x')));
drop(bridge);
}
#[test]
fn adaptive_poll_uses_idle_timeout_then_zero_timeout_to_drain_burst() {
let observed_timeouts = Arc::new(Mutex::new(Vec::new()));
let steps = VecDeque::from([
Ok(None),
Ok(Some(key_event('a'))),
Ok(Some(key_event('b'))),
Ok(None),
Err(io::Error::other("done")),
]);
let bridge = TerminalInputBridge::spawn_with_reader(ScriptedReader::new(
steps,
Arc::clone(&observed_timeouts),
));
if let Some(handle) = bridge.handle.as_ref() {
while !handle.is_finished() {
thread::sleep(Duration::from_millis(1));
}
}
let timeouts = observed_timeouts.lock().unwrap().clone();
assert_eq!(
timeouts,
vec![
INPUT_POLL_IDLE_TIMEOUT,
INPUT_POLL_IDLE_TIMEOUT,
INPUT_POLL_DRAIN_TIMEOUT,
INPUT_POLL_DRAIN_TIMEOUT,
INPUT_POLL_IDLE_TIMEOUT,
]
);
assert_eq!(bridge.receiver.try_iter().collect::<Vec<_>>().len(), 2);
drop(bridge);
}
}