Skip to content

Commit 2f79de5

Browse files
committed
cli: make burn terminal bidirectional
1 parent f7babc2 commit 2f79de5

3 files changed

Lines changed: 239 additions & 16 deletions

File tree

src/defib/cli/app.py

Lines changed: 8 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -400,24 +400,16 @@ def on_sigint(*_: object) -> None:
400400
if output == "human":
401401
console.print("[dim]--- Terminal mode (Ctrl-C to exit) ---[/dim]")
402402

403-
stop = False
404-
405-
def on_sigint(*_: object) -> None:
406-
nonlocal stop
407-
stop = True
408-
409-
signal.signal(signal.SIGINT, on_sigint)
410-
403+
from defib.cli.terminal import run_raw_terminal
411404
try:
412-
while not stop:
413-
try:
414-
data = await transport.read(256, timeout=0.1)
415-
_sys.stdout.buffer.write(data)
416-
_sys.stdout.buffer.flush()
417-
except Exception:
418-
pass
405+
await run_raw_terminal(
406+
transport,
407+
_sys.stdin.buffer,
408+
_sys.stdout.buffer,
409+
)
410+
except KeyboardInterrupt:
411+
pass
419412
finally:
420-
signal.signal(signal.SIGINT, signal.SIG_DFL)
421413
if output == "human":
422414
console.print("\n[dim]--- Terminal closed ---[/dim]")
423415

src/defib/cli/terminal.py

Lines changed: 151 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
1+
"""Bidirectional raw terminal bridge for the burn command."""
2+
3+
from __future__ import annotations
4+
5+
import asyncio
6+
import importlib
7+
import os
8+
import signal
9+
from collections.abc import Iterator
10+
from contextlib import contextmanager
11+
from typing import BinaryIO
12+
13+
from defib.transport.base import Transport, TransportTimeout
14+
15+
16+
@contextmanager
17+
def _raw_terminal(fd: int) -> Iterator[None]:
18+
"""Disable canonical input, translations, and local echo."""
19+
if os.name != "posix" or not os.isatty(fd):
20+
yield
21+
return
22+
23+
import termios
24+
import tty
25+
26+
previous = termios.tcgetattr(fd)
27+
tty.setraw(fd)
28+
try:
29+
yield
30+
finally:
31+
termios.tcsetattr(fd, termios.TCSADRAIN, previous)
32+
33+
34+
async def _pump_posix_stdin(
35+
transport: Transport,
36+
fd: int,
37+
stop: asyncio.Event,
38+
) -> None:
39+
loop = asyncio.get_running_loop()
40+
queue: asyncio.Queue[bytes] = asyncio.Queue()
41+
42+
def on_readable() -> None:
43+
try:
44+
data = os.read(fd, 1024)
45+
except OSError:
46+
data = b""
47+
if not data:
48+
loop.remove_reader(fd)
49+
queue.put_nowait(data)
50+
51+
loop.add_reader(fd, on_readable)
52+
try:
53+
while not stop.is_set():
54+
try:
55+
data = await asyncio.wait_for(queue.get(), timeout=0.1)
56+
except TimeoutError:
57+
continue
58+
if not data:
59+
stop.set()
60+
return
61+
if b"\x03" in data:
62+
before_sigint, _, _ = data.partition(b"\x03")
63+
if before_sigint:
64+
await transport.write(before_sigint)
65+
stop.set()
66+
return
67+
await transport.write(data)
68+
finally:
69+
loop.remove_reader(fd)
70+
71+
72+
async def _pump_windows_stdin(
73+
transport: Transport,
74+
stop: asyncio.Event,
75+
) -> None:
76+
msvcrt = importlib.import_module("msvcrt")
77+
78+
while not stop.is_set():
79+
if not msvcrt.kbhit():
80+
await asyncio.sleep(0.01)
81+
continue
82+
char = msvcrt.getwch()
83+
if char == "\x03":
84+
stop.set()
85+
return
86+
if char in ("\x00", "\xe0"):
87+
msvcrt.getwch()
88+
continue
89+
await transport.write(char.encode())
90+
91+
92+
async def _pump_transport(
93+
transport: Transport,
94+
stdout: BinaryIO,
95+
stop: asyncio.Event,
96+
) -> None:
97+
while not stop.is_set():
98+
try:
99+
data = await transport.read(256, timeout=0.1)
100+
except TransportTimeout:
101+
continue
102+
if not data:
103+
stop.set()
104+
return
105+
stdout.write(data)
106+
stdout.flush()
107+
108+
109+
async def _bridge_terminal(
110+
transport: Transport,
111+
stdin: BinaryIO,
112+
stdout: BinaryIO,
113+
stop: asyncio.Event,
114+
) -> None:
115+
if os.name == "nt":
116+
stdin_task = asyncio.create_task(_pump_windows_stdin(transport, stop))
117+
else:
118+
stdin_task = asyncio.create_task(_pump_posix_stdin(transport, stdin.fileno(), stop))
119+
output_task = asyncio.create_task(_pump_transport(transport, stdout, stop))
120+
stop_task = asyncio.create_task(stop.wait())
121+
tasks = {stdin_task, output_task, stop_task}
122+
123+
done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
124+
stop.set()
125+
for task in pending:
126+
task.cancel()
127+
await asyncio.gather(*pending, return_exceptions=True)
128+
129+
for task in done:
130+
if task is not stop_task:
131+
task.result()
132+
133+
134+
async def run_raw_terminal(
135+
transport: Transport,
136+
stdin: BinaryIO,
137+
stdout: BinaryIO,
138+
) -> None:
139+
"""Bridge stdin and transport until EOF, disconnect, or Ctrl-C."""
140+
stop = asyncio.Event()
141+
previous_handler = signal.getsignal(signal.SIGINT)
142+
143+
def on_sigint(*_: object) -> None:
144+
stop.set()
145+
146+
signal.signal(signal.SIGINT, on_sigint)
147+
try:
148+
with _raw_terminal(stdin.fileno()):
149+
await _bridge_terminal(transport, stdin, stdout, stop)
150+
finally:
151+
signal.signal(signal.SIGINT, previous_handler)

tests/test_terminal.py

Lines changed: 80 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,80 @@
1+
"""Tests for the interactive raw terminal bridge."""
2+
3+
import asyncio
4+
import io
5+
import os
6+
import signal
7+
8+
import pytest
9+
10+
from defib.cli.terminal import run_raw_terminal
11+
from defib.transport.mock import MockTransport
12+
13+
14+
pytestmark = pytest.mark.skipif(os.name != "posix", reason="PTY tests require POSIX")
15+
16+
17+
async def _wait_until(predicate, timeout: float = 1.0) -> None:
18+
loop = asyncio.get_running_loop()
19+
deadline = loop.time() + timeout
20+
while loop.time() < deadline:
21+
if predicate():
22+
return
23+
await asyncio.sleep(0.01)
24+
raise AssertionError("condition not met before timeout")
25+
26+
27+
@pytest.mark.asyncio
28+
async def test_pty_input_is_forwarded_and_transport_output_is_printed() -> None:
29+
import pty
30+
import termios
31+
32+
master_fd, slave_fd = pty.openpty()
33+
stdin = os.fdopen(slave_fd, "rb", buffering=0)
34+
stdout = io.BytesIO()
35+
transport = MockTransport()
36+
transport.enqueue_rx(b"hisilicon # ")
37+
38+
task = asyncio.create_task(run_raw_terminal(transport, stdin, stdout))
39+
try:
40+
await _wait_until(lambda: not termios.tcgetattr(stdin.fileno())[3] & termios.ECHO)
41+
os.write(master_fd, b"help\r")
42+
await _wait_until(lambda: b"help\r" in transport.all_tx_data)
43+
await _wait_until(lambda: b"hisilicon # " in stdout.getvalue())
44+
os.write(master_fd, b"\x03")
45+
await asyncio.wait_for(task, timeout=1.0)
46+
finally:
47+
if not task.done():
48+
task.cancel()
49+
await asyncio.gather(task, return_exceptions=True)
50+
os.close(master_fd)
51+
stdin.close()
52+
53+
54+
@pytest.mark.asyncio
55+
async def test_sigint_stops_bridge_and_restores_pty() -> None:
56+
import pty
57+
import termios
58+
59+
master_fd, slave_fd = pty.openpty()
60+
stdin = os.fdopen(slave_fd, "rb", buffering=0)
61+
stdout = io.BytesIO()
62+
transport = MockTransport()
63+
original_attrs = termios.tcgetattr(stdin.fileno())
64+
original_handler = signal.getsignal(signal.SIGINT)
65+
66+
task = asyncio.create_task(run_raw_terminal(transport, stdin, stdout))
67+
try:
68+
await _wait_until(lambda: not termios.tcgetattr(stdin.fileno())[3] & termios.ECHO)
69+
70+
signal.raise_signal(signal.SIGINT)
71+
await asyncio.wait_for(task, timeout=1.0)
72+
73+
assert termios.tcgetattr(stdin.fileno()) == original_attrs
74+
assert signal.getsignal(signal.SIGINT) == original_handler
75+
finally:
76+
if not task.done():
77+
task.cancel()
78+
await asyncio.gather(task, return_exceptions=True)
79+
os.close(master_fd)
80+
stdin.close()

0 commit comments

Comments
 (0)