From ee0da5b594e8e3824efae1f12813a72334a0c127 Mon Sep 17 00:00:00 2001 From: GuDongBin Date: Thu, 27 Aug 2026 10:22:12 +0800 Subject: [PATCH] cli: make burn terminal bidirectional --- src/defib/cli/app.py | 24 ++--- src/defib/cli/terminal.py | 179 ++++++++++++++++++++++++++++++++++++++ tests/test_terminal.py | 118 +++++++++++++++++++++++++ 3 files changed, 305 insertions(+), 16 deletions(-) create mode 100644 src/defib/cli/terminal.py create mode 100644 tests/test_terminal.py diff --git a/src/defib/cli/app.py b/src/defib/cli/app.py index 2c2c4cf..795d1a8 100644 --- a/src/defib/cli/app.py +++ b/src/defib/cli/app.py @@ -400,24 +400,16 @@ def on_sigint(*_: object) -> None: if output == "human": console.print("[dim]--- Terminal mode (Ctrl-C to exit) ---[/dim]") - stop = False - - def on_sigint(*_: object) -> None: - nonlocal stop - stop = True - - signal.signal(signal.SIGINT, on_sigint) - + from defib.cli.terminal import run_raw_terminal try: - while not stop: - try: - data = await transport.read(256, timeout=0.1) - _sys.stdout.buffer.write(data) - _sys.stdout.buffer.flush() - except Exception: - pass + await run_raw_terminal( + transport, + _sys.stdin.buffer, + _sys.stdout.buffer, + ) + except KeyboardInterrupt: + pass finally: - signal.signal(signal.SIGINT, signal.SIG_DFL) if output == "human": console.print("\n[dim]--- Terminal closed ---[/dim]") diff --git a/src/defib/cli/terminal.py b/src/defib/cli/terminal.py new file mode 100644 index 0000000..c312901 --- /dev/null +++ b/src/defib/cli/terminal.py @@ -0,0 +1,179 @@ +"""Bidirectional raw terminal bridge for the burn command.""" + +from __future__ import annotations + +import asyncio +import importlib +import os +import signal +from collections.abc import Iterator +from contextlib import contextmanager +from typing import BinaryIO + +from defib.transport.base import Transport, TransportError, TransportTimeout + + +@contextmanager +def _raw_terminal(stdin: BinaryIO) -> Iterator[None]: + """Disable canonical input, translations, and local echo.""" + if os.name != "posix" or not stdin.isatty(): + yield + return + + import termios + import tty + + fd = stdin.fileno() + previous = termios.tcgetattr(fd) + tty.setraw(fd) + try: + yield + finally: + termios.tcsetattr(fd, termios.TCSADRAIN, previous) + + +async def _pump_posix_stdin( + transport: Transport, + fd: int, + stop: asyncio.Event, +) -> None: + loop = asyncio.get_running_loop() + queue: asyncio.Queue[bytes] = asyncio.Queue() + + def on_readable() -> None: + try: + data = os.read(fd, 1024) + except OSError: + data = b"" + if not data: + loop.remove_reader(fd) + queue.put_nowait(data) + + loop.add_reader(fd, on_readable) + try: + while not stop.is_set(): + try: + data = await asyncio.wait_for(queue.get(), timeout=0.1) + except TimeoutError: + continue + if not data: + stop.set() + return + if not await _forward_input(transport, data, stop): + return + finally: + loop.remove_reader(fd) + + +async def _forward_input( + transport: Transport, + data: bytes, + stop: asyncio.Event, +) -> bool: + """Forward input bytes, treating Ctrl-C as a local terminal command.""" + before_sigint, separator, _ = data.partition(b"\x03") + if before_sigint: + await transport.write(before_sigint) + if separator: + stop.set() + return False + return True + + +async def _pump_stream_stdin( + transport: Transport, + stdin: BinaryIO, + stop: asyncio.Event, +) -> None: + while not stop.is_set(): + data = await asyncio.to_thread(stdin.read, 1024) + if not data: + stop.set() + return + if not await _forward_input(transport, data, stop): + return + + +async def _pump_windows_console( + transport: Transport, + stop: asyncio.Event, +) -> None: + msvcrt = importlib.import_module("msvcrt") + + while not stop.is_set(): + if not msvcrt.kbhit(): + await asyncio.sleep(0.01) + continue + char = msvcrt.getch() + if char in (b"\x00", b"\xe0"): + msvcrt.getch() + continue + if not await _forward_input(transport, char, stop): + return + + +async def _pump_transport( + transport: Transport, + stdout: BinaryIO, + stop: asyncio.Event, +) -> None: + while not stop.is_set(): + try: + data = await transport.read(256, timeout=0.1) + except TransportTimeout: + continue + if not data: + stop.set() + return + stdout.write(data) + stdout.flush() + + +async def _bridge_terminal( + transport: Transport, + stdin: BinaryIO, + stdout: BinaryIO, + stop: asyncio.Event, +) -> None: + if os.name == "nt" and stdin.isatty(): + stdin_task = asyncio.create_task(_pump_windows_console(transport, stop)) + elif stdin.isatty(): + stdin_task = asyncio.create_task(_pump_posix_stdin(transport, stdin.fileno(), stop)) + else: + stdin_task = asyncio.create_task(_pump_stream_stdin(transport, stdin, stop)) + output_task = asyncio.create_task(_pump_transport(transport, stdout, stop)) + stop_task = asyncio.create_task(stop.wait()) + tasks = {stdin_task, output_task, stop_task} + + done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED) + stop.set() + for task in pending: + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + + for task in done: + if task is not stop_task: + task.result() + + +async def run_raw_terminal( + transport: Transport, + stdin: BinaryIO, + stdout: BinaryIO, +) -> None: + """Bridge stdin and transport until EOF, disconnect, or Ctrl-C.""" + stop = asyncio.Event() + previous_handler = signal.getsignal(signal.SIGINT) + + def on_sigint(*_: object) -> None: + stop.set() + + signal.signal(signal.SIGINT, on_sigint) + try: + with _raw_terminal(stdin): + try: + await _bridge_terminal(transport, stdin, stdout, stop) + except TransportError: + pass + finally: + signal.signal(signal.SIGINT, previous_handler) diff --git a/tests/test_terminal.py b/tests/test_terminal.py new file mode 100644 index 0000000..46020a8 --- /dev/null +++ b/tests/test_terminal.py @@ -0,0 +1,118 @@ +"""Tests for the interactive raw terminal bridge.""" + +import asyncio +import io +import os +import signal + +import pytest + +from defib.cli.terminal import run_raw_terminal +from defib.transport.base import TransportError +from defib.transport.mock import MockTransport + + +requires_pty = pytest.mark.skipif(os.name != "posix", reason="PTY tests require POSIX") + + +async def _wait_until(predicate, timeout: float = 1.0) -> None: + loop = asyncio.get_running_loop() + deadline = loop.time() + timeout + while loop.time() < deadline: + if predicate(): + return + await asyncio.sleep(0.01) + raise AssertionError("condition not met before timeout") + + +@pytest.mark.asyncio +@requires_pty +async def test_pty_input_is_forwarded_and_transport_output_is_printed() -> None: + import pty + import termios + + master_fd, slave_fd = pty.openpty() + stdin = os.fdopen(slave_fd, "rb", buffering=0) + stdout = io.BytesIO() + transport = MockTransport() + transport.enqueue_rx(b"hisilicon # ") + + task = asyncio.create_task(run_raw_terminal(transport, stdin, stdout)) + try: + await _wait_until(lambda: not termios.tcgetattr(stdin.fileno())[3] & termios.ECHO) + os.write(master_fd, b"help\r") + await _wait_until(lambda: b"help\r" in transport.all_tx_data) + await _wait_until(lambda: b"hisilicon # " in stdout.getvalue()) + os.write(master_fd, b"\x03") + await asyncio.wait_for(task, timeout=1.0) + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + os.close(master_fd) + stdin.close() + + +@pytest.mark.asyncio +@requires_pty +async def test_sigint_stops_bridge_and_restores_pty() -> None: + import pty + import termios + + master_fd, slave_fd = pty.openpty() + stdin = os.fdopen(slave_fd, "rb", buffering=0) + stdout = io.BytesIO() + transport = MockTransport() + original_attrs = termios.tcgetattr(stdin.fileno()) + original_handler = signal.getsignal(signal.SIGINT) + + task = asyncio.create_task(run_raw_terminal(transport, stdin, stdout)) + try: + await _wait_until(lambda: not termios.tcgetattr(stdin.fileno())[3] & termios.ECHO) + + signal.raise_signal(signal.SIGINT) + await asyncio.wait_for(task, timeout=1.0) + + assert termios.tcgetattr(stdin.fileno()) == original_attrs + assert signal.getsignal(signal.SIGINT) == original_handler + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + os.close(master_fd) + stdin.close() + + +@pytest.mark.asyncio +async def test_redirected_stdin_is_forwarded_until_eof() -> None: + stdin = io.BytesIO(b"help\r") + stdout = io.BytesIO() + transport = MockTransport() + + await run_raw_terminal(transport, stdin, stdout) + + assert transport.all_tx_data == b"help\r" + + +@pytest.mark.asyncio +@requires_pty +async def test_transport_disconnect_ends_terminal_cleanly() -> None: + import pty + + class DisconnectingTransport(MockTransport): + async def read(self, size: int, timeout: float | None = None) -> bytes: + raise TransportError("remote disconnected") + + master_fd, slave_fd = pty.openpty() + stdin = os.fdopen(slave_fd, "rb", buffering=0) + task = asyncio.create_task( + run_raw_terminal(DisconnectingTransport(), stdin, io.BytesIO()) + ) + try: + await asyncio.wait_for(task, timeout=1.0) + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + os.close(master_fd) + stdin.close()