tmux-amq/src/tamq/ptytap.py
tegwick 6d2ccc7760
Some checks failed
tamq-ci / test (push) Failing after 5s
feat: complete reliable coordination adapter
Assistant: codex
Assistant-Model: gpt-5.6-sol
Assistant-Session: 01a03397-4d51-7fd1-8ff2-946eb22ea2bc
2026-08-26 08:11:09 +02:00

332 lines
12 KiB
Python

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:]