from __future__ import annotations import fcntl import os import pty import re import select import signal import sys import termios import tty import codecs import time from collections import deque from collections.abc import Callable, Sequence from pathlib import Path from .broker import InputBroker from .terminal import terminal_safe ENHANCED_ENTER = re.compile( rb"\x1b(?:\[13(?:;[0-9:]+)*u|\[27;[0-9:]+;13~|OM)" ) TERMINAL_INPUT_SEQUENCE = re.compile( rb"\x1b(?:" rb"\[[0-?]*[ -/]*[@-~]" rb"|O." rb"|P.*?(?:\x1b\\)" rb"|\].*?(?:\x07|\x1b\\)" rb")", re.DOTALL, ) def observed_line(raw: bytes) -> str: """Return typed text without terminal replies or local editing controls.""" decoded = TERMINAL_INPUT_SEQUENCE.sub(b"", raw).decode( "utf-8", errors="replace" ) text: list[str] = [] for character in decoded: if character in ("\b", "\x7f"): if text: text.pop() elif character == "\t" or character >= " ": text.append(character) return "".join(text) class TerminalOutputObserver: """Normalize PTY output into logical lines without interpreting its content.""" _CURSOR_BOUNDARIES = frozenset("ABCDEFGHfJKd") def __init__(self, on_line: Callable[[str], None], *, echo_seconds: float = 5.0): self.on_line = on_line self.echo_seconds = echo_seconds self.decoder = codecs.getincrementaldecoder("utf-8")("replace") self.text: list[str] = [] self.state = "normal" self.string_kind = "" self.csi: list[str] = [] self.cursor_row: int | None = None self.after_cr = False self.empty_cr = False self.operator_lines: deque[tuple[str, float]] = deque(maxlen=32) self.last_redraw: str | None = None def note_operator_line(self, line: str) -> None: self.operator_lines.append((line, time.monotonic() + self.echo_seconds)) def _is_recent_operator_echo(self, line: str) -> bool: now = time.monotonic() while self.operator_lines and self.operator_lines[0][1] < now: self.operator_lines.popleft() return any(candidate == line for candidate, _ in self.operator_lines) def _emit(self, *, redraw: bool = False, allow_empty: bool = False) -> None: line = "".join(self.text) self.text.clear() if not line: if allow_empty: self.on_line("") return if self._is_recent_operator_echo(line): return if redraw and line == self.last_redraw: return self.last_redraw = line if redraw else None self.on_line(line) def _move_cursor(self, final: str) -> None: parameters = "".join(self.csi) previous = self.cursor_row new_row: int | None = None try: first = int((parameters.split(";", 1)[0] or "1")) except ValueError: first = 1 if final in "Hf": new_row = first elif final == "d": new_row = first elif final in "BE" and previous is not None: new_row = previous + first elif final in "AF" and previous is not None: new_row = max(1, previous - first) if ( previous is not None and new_row is not None and new_row > previous + 1 ): self.on_line("") if new_row is not None: self.cursor_row = new_row def feed(self, data: bytes) -> None: for character in self.decoder.decode(data): if self.state == "normal": if character == "\x1b": self.after_cr = False self.state = "esc" elif character == "\r": self.empty_cr = not self.text self._emit(redraw=True) self.after_cr = True elif character == "\n": if self.after_cr: if self.empty_cr: self._emit(allow_empty=True) else: self._emit(allow_empty=True) self.after_cr = False self.empty_cr = False if self.cursor_row is not None: self.cursor_row += 1 elif character in ("\b", "\x7f"): self.after_cr = False if self.text: self.text.pop() elif character == "\t" or character >= " ": self.after_cr = False self.text.append(character) elif self.state == "esc": if character == "[": self.state = "csi" self.csi.clear() elif character in "]P_^": self.state = "string" self.string_kind = character else: self.state = "normal" elif self.state == "csi": if "@" <= character <= "~": if character in self._CURSOR_BOUNDARIES: self._emit(redraw=True) self._move_cursor(character) self.state = "normal" else: self.csi.append(character) elif self.state == "string": if character == "\x07" and self.string_kind == "]": self.state = "normal" elif character == "\x1b": self.state = "string_esc" elif self.state == "string_esc": self.state = "normal" if character == "\\" else "string" def flush(self) -> None: self.decoder.decode(b"", final=True) self._emit() class PtyTap: """Full-duplex PTY proxy. Input is observed, never rewritten.""" def __init__(self, command: Sequence[str], broker: InputBroker, *, on_line: Callable[[str], None] | None = None, ready_file: Path | None = None): self.command = list(command) self.broker = broker self.on_line = on_line self.ready_file = ready_file self.output_observer = TerminalOutputObserver(self._observe_worker_line) def _observe_worker_line(self, line: str) -> None: inspector = getattr(self.broker, "inspect_worker_line", None) routed = inspector(line) if inspector is not None else None if routed and self.on_line: self.on_line(line) def _local_notice(self, line: str) -> None: payload = f"\x1b7\r\n{terminal_safe(line)}\r\n\x1b8".encode("utf-8") write_all(sys.stdout.fileno(), payload) def _observe_input(self, buffer: bytearray, data: bytes) -> None: buffer.extend(data) while True: boundaries = [ (index, index + 1) for marker in (b"\r", b"\n") if (index := buffer.find(marker)) >= 0 ] if match := ENHANCED_ENTER.search(buffer): boundaries.append(match.span()) if not boundaries: return start, end = min(boundaries) raw = bytes(buffer[:start]) del buffer[:end] line = observed_line(raw) self.output_observer.note_operator_line(line) inspector = getattr(self.broker, "inspect_operator_line", self.broker.inspect_line) routed = inspector(line) if routed and self.on_line: self.on_line(line) def run(self) -> int: stdin_fd = sys.stdin.fileno() stdout_fd = sys.stdout.fileno() master, slave = pty.openpty() copy_winsize(stdin_fd, slave) pid = os.fork() if pid == 0: try: os.close(master) os.setsid() fcntl.ioctl(slave, termios.TIOCSCTTY, 0) for fd in (0, 1, 2): os.dup2(slave, fd) if slave > 2: os.close(slave) os.execvp(self.command[0], self.command) except BaseException as exc: os.write(2, f"tamq tap: cannot start {self.command[0]}: {exc}\n".encode()) os._exit(127) os.close(slave) input_buffer = bytearray() saved_terminal = termios.tcgetattr(stdin_fd) if os.isatty(stdin_fd) else None saved_handlers: dict[int, signal.Handlers] = {} input_open = True def resize(_signum=None, _frame=None) -> None: copy_winsize(stdin_fd, master) def forward(signum, _frame) -> None: try: os.killpg(pid, signum) except ProcessLookupError: pass status = 0 try: self.broker.notify = self._local_notice if self.ready_file is not None: self.ready_file.parent.mkdir(parents=True, exist_ok=True) self.ready_file.write_text(str(os.getpid()), encoding="utf-8") if saved_terminal is not None: tty.setraw(stdin_fd) saved_handlers[signal.SIGWINCH] = signal.signal(signal.SIGWINCH, resize) for signum in (signal.SIGINT, signal.SIGTERM, signal.SIGHUP): saved_handlers[signum] = signal.signal(signum, forward) resize() while True: sources = [master] if input_open: sources.append(stdin_fd) try: readable, _, _ = select.select(sources, [], []) except InterruptedError: continue if master in readable: try: data = os.read(master, 65536) except OSError: break if not data: break write_all(stdout_fd, data) self.output_observer.feed(data) if input_open and stdin_fd in readable: data = os.read(stdin_fd, 65536) if not data: input_open = False try: os.killpg(pid, signal.SIGHUP) except ProcessLookupError: pass continue write_all(master, data) self._observe_input(input_buffer, data) except BaseException: try: os.killpg(pid, signal.SIGHUP) except ProcessLookupError: pass raise finally: if self.ready_file is not None: self.ready_file.unlink(missing_ok=True) self.output_observer.flush() flush_worker = getattr(self.broker, "flush_worker_message", None) if flush_worker is not None: flush_worker("worker_exit") for signum, handler in saved_handlers.items(): signal.signal(signum, handler) if saved_terminal is not None: termios.tcsetattr(stdin_fd, termios.TCSAFLUSH, saved_terminal) os.close(master) try: _, status = os.waitpid(pid, 0) except ChildProcessError: status = 0 return os.waitstatus_to_exitcode(status) def copy_winsize(source_fd: int, target_fd: int) -> bool: """Copy terminal rows/columns and pixel dimensions when a source tty exists.""" try: size = fcntl.ioctl(source_fd, termios.TIOCGWINSZ, b"\0" * 8) fcntl.ioctl(target_fd, termios.TIOCSWINSZ, size) return True except OSError: return False def write_all(fd: int, data: bytes) -> None: view = memoryview(data) while view: written = os.write(fd, view) if written == 0: raise OSError("terminal write returned zero bytes") view = view[written:]