Fix ttyfd passing

This commit is contained in:
Kovid Goyal 2022-07-07 20:30:04 +05:30
parent 72f3e8cd40
commit 991fbacb99
No known key found for this signature in database
GPG key ID: 06BC317B515ACE7C
2 changed files with 21 additions and 7 deletions

View file

@ -392,7 +392,7 @@ def __init__(self, conn: socket.socket, addr: bytes, poll: select.poll):
self.cwd = ''
self.env: Dict[str, str] = {}
self.argv: List[str] = []
self.stdin = self.stdout = self.stderr = -1
self.stdin = self.stdout = self.stderr = self.tty_fd = -1
self.pid = -1
self.closed = False
@ -431,12 +431,14 @@ def read(self) -> bool:
for x in self.fds:
os.set_inheritable(x, x is not self.fds[0])
os.set_blocking(x, True)
self.tty_fd = self.fds[0]
if self.stdin > -1:
self.stdin = self.fds[self.stdin]
if self.stdout > -1:
self.stdout = self.fds[self.stdout]
if self.stderr > -1:
self.stderr = self.fds[self.stderr]
del self.fds[:]
return True
elif cmd == 'cwd':
self.cwd = payload
@ -481,16 +483,17 @@ def fork(self, free_non_child_resources: Callable[[], None]) -> None:
os.close(r)
os.setsid()
restore_python_signal_handlers()
if self.fds:
if self.tty_fd > -1:
sys.__stdout__.flush()
sys.__stderr__.flush()
establish_controlling_tty(
self.fds[0],
self.tty_fd,
sys.__stdin__.fileno() if self.stdin < 0 else -1,
sys.__stdout__.fileno() if self.stdout < 0 else -1,
sys.__stderr__.fileno() if self.stderr < 0 else -1)
# the std streams fds are in all_non_child_fds already
# so they will be closed there
self.tty_fd = -1
# the std streams fds are closed in free_non_child_resources(), see
# SocketChild.close()
if self.stdin > -1:
os.dup2(self.stdin, sys.__stdin__.fileno())
if self.stdout > -1:
@ -531,6 +534,12 @@ def close(self) -> None:
self.unregister_from_poll()
self.closed = True
self.conn.close()
for x in self.fds:
os.close(x)
del self.fds[:]
if self.tty_fd > -1:
os.close(self.tty_fd)
self.tty_fd = -1
if self.stdin > -1:
os.close(self.stdin)
self.stdin = -1

View file

@ -484,7 +484,12 @@ struct pollees {
static struct pollees pollees = {0};
#define register_for_poll(which) pollees.poll_data[pollees.num_fds].fd = which; pollees.poll_data[pollees.num_fds].events = POLLIN; pollees.idx.which = pollees.num_fds++;
#define unregister_for_poll(which) if (pollees.idx.which > -1) { remove_i_from_array(pollees.poll_data, pollees.idx.which, pollees.num_fds); pollees.idx.which = -1; }
#define unregister_for_poll(which) if (pollees.idx.which > -1) { remove_i_from_array(pollees.poll_data, pollees.idx.which, pollees.num_fds); \
if (pollees.idx.self_ttyfd > pollees.idx.which) pollees.idx.self_ttyfd--; \
if (pollees.idx.signal_read_fd > pollees.idx.which) pollees.idx.signal_read_fd--; \
if (pollees.idx.socket_fd > pollees.idx.which) pollees.idx.socket_fd--; \
if (pollees.idx.child_master_fd > pollees.idx.which) pollees.idx.child_master_fd--; \
pollees.idx.which = -1; }
#define set_poll_events(which, val) if (pollees.idx.which > -1) { pollees.poll_data[pollees.idx.which].events = (val); }
#define poll_revents(which) ((pollees.idx.which > -1) ? pollees.poll_data[pollees.idx.which].revents : 0)
@ -492,7 +497,7 @@ static void
loop(void) {
#define fail(s) { print_error(s, errno); return; }
#define check_fd(name) { if (poll_revents(name) & POLLERR) { pe("File descriptor %s failed", #name); return; } if (poll_revents(name) & POLLHUP) { pe("File descriptor %s hungup", #name); return; } }
register_for_poll(self_ttyfd); register_for_poll(socket_fd); register_for_poll(signal_read_fd); register_for_poll(child_master_fd);
register_for_poll(self_ttyfd); register_for_poll(signal_read_fd); register_for_poll(socket_fd); register_for_poll(child_master_fd);
while (true) {
int ret;