Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion projects/fal/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ readme = "README.md"
requires-python = ">=3.8"
dependencies = [
"isolate[build]>=0.26.12,<0.27.0",
"isolate-proto>=0.34.2,<1",
"isolate-proto>=0.34.3,<1",
"grpcio>=1.64.0,<2",
"dill==0.3.7",
"cloudpickle>=3.0.0,<3.2",
Expand Down
49 changes: 40 additions & 9 deletions projects/fal/src/fal/cli/runners.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

import grpc
import httpx
import isolate_proto
from httpx_sse import connect_sse
from rich.console import Console
from structlog.typing import EventDict
Expand Down Expand Up @@ -133,7 +134,8 @@ def _get_tty_size(fd: int):

def _shell(args):
"""Open an interactive shell on a runner."""
return _shell_session(args, command=None, interactive=True)
# Always a PTY: the login shell expects one even when local stdin is piped.
return _shell_session(args, command=None, interactive=True, remote_tty=True)


def _exec(args):
Expand All @@ -147,12 +149,27 @@ def _exec(args):
args.console.print("[red]Error:[/] No command specified.")
return 1

return _shell_session(args, command=command, interactive=args.interactive)
# A PTY mangles bytes (echo, CR/NL translation, signal characters), so only
# ask for one when a real terminal is attached.
remote_tty = args.interactive and sys.stdin.isatty()
return _shell_session(
args, command=command, interactive=args.interactive, remote_tty=remote_tty
)


def _shell_session(args, command, interactive, *, remote_tty, stdout=None, stderr=None):
"""Stream a shell session on a runner; command=None opens a login shell.

def _shell_session(args, command, interactive):
"""Stream a shell session on a runner; command=None opens a login shell."""
import isolate_proto
`remote_tty` controls whether the remote command runs under a pseudo-terminal.
Without one, remote stdin is closed once local input ends so pipe-reading
commands see EOF. Remote output goes to `stdout` and `stderr` (default: this
process's corresponding streams). Servers predating stream identification
produce output with no stream set, which is treated as stdout.
"""
if stdout is None:
stdout = sys.stdout.buffer
if stderr is None:
stderr = sys.stderr.buffer

client = SyncServerlessClient(host=args.host, team=args.team)
stub = client._create_host()._connection.stub
Expand Down Expand Up @@ -204,9 +221,13 @@ def read_stdin():
def stream_inputs():
"""Generate input stream for gRPC."""
# Send initial message with runner_id
yield isolate_proto.ShellRunnerInput(runner_id=runner_id, command=command)
yield isolate_proto.ShellRunnerInput(
runner_id=runner_id, command=command, tty=remote_tty
)

if not interactive:
if not remote_tty:
yield isolate_proto.ShellRunnerInput(close=True)
return

# Send terminal size
Expand All @@ -233,6 +254,8 @@ def stream_inputs():
msg.tty_size.width = w
yield msg
elif msg_type == "eof":
if not remote_tty:
yield isolate_proto.ShellRunnerInput(close=True)
return

exit_code = 1
Expand All @@ -247,8 +270,13 @@ def restore_tty() -> None:
exit_code = output.exit_code
break
if output.data:
sys.stdout.buffer.write(output.data)
sys.stdout.buffer.flush()
stream = (
stderr
if output.HasField("stream") and output.stream == 2
else stdout
)
stream.write(output.data)
stream.flush()
if output.close:
break
exit_code = exit_code or 0
Expand Down Expand Up @@ -854,7 +882,10 @@ def _add_exec_parser(subparsers, parents):
"-it",
"--interactive",
action="store_true",
help="Allocate a TTY and attach stdin (interactive mode).",
help=(
"Attach stdin. A TTY is allocated only when stdin is a terminal; "
"with piped stdin the command gets a raw byte stream."
),
)
# PARSER keeps fal's own flags parseable between the runner id and the
# command; REMAINDER would swallow them into the command.
Expand Down
132 changes: 131 additions & 1 deletion projects/fal/tests/unit/cli/test_runners.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,15 @@
import io
import json
import os
import struct
from unittest.mock import MagicMock, patch

import isolate_proto
import pytest

from fal.cli.main import parse_args
from fal.cli.parser import FalParserExit
from fal.cli.runners import _exec, _get_tty_size, _gpus
from fal.cli.runners import _exec, _get_tty_size, _gpus, _shell_session

if os.name != "nt":
import fcntl
Expand Down Expand Up @@ -129,6 +131,134 @@ def test_exec_rejects_separator_only_command(mock_client_cls):
assert "No command specified" in console.print.call_args[0][0]


@patch("fal.cli.runners.SyncServerlessClient")
def test_shell_session_separates_output_streams(mock_client_cls):
stub = mock_client_cls.return_value._create_host.return_value._connection.stub
stub.ShellRunner.return_value = [
isolate_proto.ShellRunnerOutput(data=b"out"),
isolate_proto.ShellRunnerOutput(data=b"err", stream=2),
isolate_proto.ShellRunnerOutput(exit_code=0),
]
stdout = io.BytesIO()
stderr = io.BytesIO()
args = parse_args(["runners", "exec", "runner-id", "--", "true"])

exit_code = _shell_session(
args,
command=["true"],
interactive=False,
remote_tty=False,
stdout=stdout,
stderr=stderr,
)

assert exit_code == 0
assert stdout.getvalue() == b"out"
assert stderr.getvalue() == b"err"


@patch("fal.cli.runners.SyncServerlessClient")
def test_exec_without_terminal_requests_no_tty_and_closes_stdin(mock_client_cls):
with patch("fal.cli.runners.sys.stdin") as stdin:
stdin.isatty.return_value = False
exit_code, sent, _ = _exec_with_command(mock_client_cls, ["python"])

assert exit_code == 0
assert sent[0].HasField("tty") and sent[0].tty is False
assert [msg.close for msg in sent] == [False, True]


@pytest.mark.skipif(os.name == "nt", reason="Interactive shell is Unix-only")
@patch("fal.cli.runners.SyncServerlessClient")
def test_exec_interactive_with_piped_stdin_streams_raw_bytes(mock_client_cls):
sent = []

def shell_runner(inputs):
sent.extend(inputs)
return [isolate_proto.ShellRunnerOutput(exit_code=0)]

stub = mock_client_cls.return_value._create_host.return_value._connection.stub
stub.ShellRunner.side_effect = shell_runner

read_fd, write_fd = os.pipe()
os.write(write_fd, b"SSH-2.0-client\r\n")
os.close(write_fd)
args = parse_args(["runners", "exec", "-it", "runner-id", "--", "sshd", "-i"])
args.console = MagicMock()
try:
with patch("fal.cli.runners.sys.stdin") as stdin:
stdin.isatty.return_value = False
stdin.fileno.return_value = read_fd
assert args.func(args) == 0
finally:
os.close(read_fd)

assert sent[0].tty is False
assert not sent[0].HasField("tty_size")
assert list(sent[0].command) == ["sshd", "-i"]
assert [msg.data for msg in sent[1:-1]] == [b"SSH-2.0-client\r\n"]
assert sent[-1].close is True


@pytest.mark.skipif(os.name == "nt", reason="Pseudo-terminals are Unix-only")
@patch("fal.cli.runners.SyncServerlessClient")
def test_exec_interactive_on_terminal_requests_tty_and_sends_size(mock_client_cls):
sent = []

def shell_runner(inputs):
for msg in inputs:
sent.append(msg)
if msg.HasField("tty_size"):
break
return [isolate_proto.ShellRunnerOutput(exit_code=0)]

stub = mock_client_cls.return_value._create_host.return_value._connection.stub
stub.ShellRunner.side_effect = shell_runner

master_fd, slave_fd = os.openpty()
args = parse_args(["runners", "exec", "-it", "runner-id", "--", "bash"])
args.console = MagicMock()
try:
with patch("fal.cli.runners.sys.stdin") as stdin:
stdin.isatty.return_value = True
stdin.fileno.return_value = slave_fd
assert args.func(args) == 0
finally:
os.close(master_fd)
os.close(slave_fd)

# The local `import tty` inside the session must not leak into the message.
assert sent[0].tty is True
assert sent[1].HasField("tty_size")


@patch("fal.cli.runners.SyncServerlessClient")
def test_shell_always_requests_tty(mock_client_cls):
sent = []

def shell_runner(inputs):
sent.extend(inputs)
return [isolate_proto.ShellRunnerOutput(exit_code=0)]

stub = mock_client_cls.return_value._create_host.return_value._connection.stub
stub.ShellRunner.side_effect = shell_runner

read_fd, write_fd = os.pipe()
os.close(write_fd)
args = parse_args(["runners", "shell", "runner-id"])
args.console = MagicMock()
try:
with patch("fal.cli.runners.sys.stdin") as stdin:
stdin.isatty.return_value = False
stdin.fileno.return_value = read_fd
assert args.func(args) == 0
finally:
os.close(read_fd)

assert sent[0].tty is True
assert all(not msg.close for msg in sent)
Comment thread
badayvedat marked this conversation as resolved.


def _mock_client(payload):
client = MagicMock()
client.runners.gpus.return_value = payload
Expand Down
Loading