magi-code 0.63.2

Repository-aware CLI coding agent for terminal work
Documentation
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,
{
    // crossterm 0.29 exposes no public wake handle for blocking read(). Its internal
    // waker is crate-private and tied to EventStream, so this uses a bounded poll:
    // idle waits cut wakeups; zero-timeout drains queued bursts without the old 10ms tick.
    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);
    }
}