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
19 changes: 17 additions & 2 deletions docs/deploy.md
Original file line number Diff line number Diff line change
Expand Up @@ -221,13 +221,28 @@ idle sandbox; a connection through the gateway wakes it (approximately
6 seconds) and keeps it awake while connected. The first data can
therefore come after a short delay.

The gateway serves an interactive shell, command execution, and SFTP.
It refuses scp (the legacy protocol) and port forwarding. SFTP runs the
The gateway serves an interactive shell, command execution, SFTP, and
local port forwarding to the sandbox loopback. It refuses scp (the
legacy protocol) and remote port forwarding. SFTP runs the
sandbox's own SFTP server over one persistent session for each SSH
connection, thus repeated file operations on one connection start no new
session. The session closes after a short idle period, so an inactive
sandbox still sleeps.

Forward a local port to a service that listens on the sandbox loopback:

```bash
ssh -N -p 2222 -i <private-key> -L 43123:127.0.0.1:43123 <host-name>@<gateway-address>
```

The destination must be `127.0.0.1` or `localhost`. The gateway refuses
all other destinations, also the IPv6 loopback. For `docker-sbx`, the
gateway opens one tunnel through `sbx ssh proxy` for each caller
connection. All forwarding channels of that caller connection use it,
and it closes when the caller disconnects. sandboxd accepts a channel to
a closed port and then closes it immediately. Thus the caller gets an
immediate end of data, not a channel-open error.

Every session runs with HOME set to the per-host home directory
`/home/<host-name>`. The gateway makes this directory and moves into it
first. Thus a caller writes files there — for example a `.gitconfig` or
Expand Down
104 changes: 65 additions & 39 deletions src/gateway/server.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,12 @@
import asyncio
import contextlib
import logging
from typing import cast

import asyncssh
import asyncssh.sftp
from asyncssh.constants import OPEN_ADMINISTRATIVELY_PROHIBITED, OPEN_CONNECT_FAILED
from asyncssh.forward import SSHForwarder
from sqlalchemy import select

from core.database import async_session_factory
Expand All @@ -19,50 +22,76 @@

_RECEIVE_CHUNK_BYTES = 32768

# The gateway forwards each file operation to the sandbox by an opaque
# handle. Two SFTP extensions ask the server to seek inside an open file:
# server-side copy and sparse-range detection. The gateway cannot serve
# them on an opaque handle. asyncssh advertises them from this class list
# and gives no per-server control, so the gateway removes them from the
# list. The filter is idempotent, thus it can run on each SFTP session.
_UNSUPPORTED_SFTP_EXTENSIONS = (b"copy-data", b"ranges@asyncssh.com")

# IPv6 loopback does not reach the sandbox, so only IPv4 names are allowed.
_LOOPBACK = frozenset({"127.0.0.1", "localhost"})

def _disable_unsupported_sftp_extensions() -> None:
asyncssh.sftp.SFTPServerHandler._extensions = [
extension
for extension in asyncssh.sftp.SFTPServerHandler._extensions
if extension[0] not in _UNSUPPORTED_SFTP_EXTENSIONS
]
# Server-side copy and sparse ranges seek inside an open file, which the
# sandbox's opaque handles cannot do. asyncssh lists them per class only.
_UNSUPPORTED_SFTP_EXTENSIONS = (b"copy-data", b"ranges@asyncssh.com")


class GatewayConnection(asyncssh.SSHServer):
"""One caller connection. The key is the identity; the username must name
the same host, so one leaked key cannot probe other host names."""

def __init__(self) -> None:
self.host: Host | None = None
# Set by validate_public_key: asyncssh opens no channel before auth.
self.host: Host
self._connection: asyncssh.SSHServerConnection
self._sftp_backend: SandboxSftpBackend | None = None
self._cleanup: asyncio.Task[None] | None = None
self._tunnel: asyncio.Future[asyncssh.SSHClientConnection] | None = None

def sftp_backend(self) -> SandboxSftpBackend:
def get_sftp_backend(self) -> SandboxSftpBackend:
"""Return the connection's one SFTP backend, shared by every SFTP
session. It is made on first use; its process opens lazily."""
assert self.host is not None
if self._sftp_backend is None:
if not self._sftp_backend:
provider = get_vm_provider(self.host.provider)
if not provider.gateway_process_class:
raise asyncssh.SFTPOpUnsupported("cannot open a session for this host")
self._sftp_backend = SandboxSftpBackend(provider.gateway_process_class, self.host.name)
return self._sftp_backend

def connection_made(self, conn: asyncssh.SSHServerConnection) -> None:
self._connection = conn

def connection_lost(self, exc: Exception | None) -> None:
# Close the backend, so its exec process ends and the sandbox can
# sleep. Keep the task until it finishes: the loop can otherwise
# collect it during the cleanup.
if self._sftp_backend:
with contextlib.suppress(RuntimeError):
self._cleanup = asyncio.get_running_loop().create_task(self._sftp_backend.aclose())
if self._tunnel:
# A tunnel that is still opening stops; an open one closes.
if not self._tunnel.done():
self._tunnel.cancel()
elif not self._tunnel.exception():
self._tunnel.result().close()

async def connection_requested( # pyright: ignore[reportIncompatibleMethodOverride]
self, dest_host: str, dest_port: int, orig_host: str, orig_port: int
) -> SSHForwarder:
"""Forward a TCP channel to the sandbox's own loopback. Every channel
of this connection shares one tunnel into the sandbox."""
if dest_host not in _LOOPBACK:
raise asyncssh.ChannelOpenError(
OPEN_ADMINISTRATIVELY_PROHIBITED, "only sandbox loopback is allowed"
)
if not self._tunnel:
provider = get_vm_provider(self.host.provider)
self._tunnel = asyncio.ensure_future(provider.open_gateway_tunnel(self.host.name))
try:
tunnel = await self._tunnel
except ProviderError as error:
self._tunnel = None
logger.warning("gateway: forward failed for host=%s: %s", self.host.name, error)
raise asyncssh.ChannelOpenError(
OPEN_CONNECT_FAILED, "cannot connect to this host"
) from error
forwarder = await self._connection.forward_tunneled_connection(tunnel, dest_host, dest_port)
logger.info("gateway: forward open host=%s port=%d", self.host.name, dest_port)
return forwarder

def begin_auth(self, username: str) -> bool:
return True
Expand Down Expand Up @@ -90,9 +119,7 @@ async def validate_public_key(self, username: str, key: asyncssh.SSHKey) -> bool


async def _bridge(process: asyncssh.SSHServerProcess) -> None:
connection = process.channel.get_connection()
server = connection.get_owner()
assert isinstance(server, GatewayConnection) and server.host is not None
server = cast(GatewayConnection, process.channel.get_connection().get_owner())
host = server.host

terminal: TerminalSize | None = None
Expand Down Expand Up @@ -170,30 +197,30 @@ async def _pump_channel_to_process(
sandbox_process.send(data)


def _load_host_key(settings: GatewaySettings) -> asyncssh.SSHKey:
path = settings.host_key_path
if path.exists():
return asyncssh.read_private_key(path)
path.parent.mkdir(parents=True, exist_ok=True)
key = asyncssh.generate_private_key("ssh-ed25519")
path.touch(mode=0o600)
path.write_bytes(key.export_private_key("openssh"))
logger.info("gateway: made a new host key at %s", path)
return key


def _open_sftp(channel: asyncssh.SSHServerChannel) -> GatewaySFTPServer:
_disable_unsupported_sftp_extensions()
server = channel.get_connection().get_owner()
assert isinstance(server, GatewayConnection)
return GatewaySFTPServer(channel, server.sftp_backend())
asyncssh.sftp.SFTPServerHandler._extensions = [
extension
for extension in asyncssh.sftp.SFTPServerHandler._extensions
if extension[0] not in _UNSUPPORTED_SFTP_EXTENSIONS
]
server = cast(GatewayConnection, channel.get_connection().get_owner())
return GatewaySFTPServer(channel, server.get_sftp_backend())


async def start(settings: GatewaySettings) -> asyncssh.SSHAcceptor:
path = settings.host_key_path
if path.exists():
host_key = asyncssh.read_private_key(path)
else:
path.parent.mkdir(parents=True, exist_ok=True)
host_key = asyncssh.generate_private_key("ssh-ed25519")
path.touch(mode=0o600)
path.write_bytes(host_key.export_private_key("openssh"))
logger.info("gateway: made a new host key at %s", path)
server = await asyncssh.listen(
host=settings.bind_host,
port=settings.ssh_port,
server_host_keys=[_load_host_key(settings)],
server_host_keys=[host_key],
server_factory=GatewayConnection,
process_factory=_bridge,
encoding=None,
Expand All @@ -212,7 +239,6 @@ async def serve() -> None:


if __name__ == "__main__":
# Service entry point: `python -m gateway.server`.
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s %(message)s",
Expand Down
116 changes: 116 additions & 0 deletions src/gateway/tests/test_server.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,13 @@
import asyncio
import hashlib
import os
from datetime import UTC, datetime
from types import SimpleNamespace
from typing import ClassVar

import asyncssh
import pytest
from asyncssh.constants import OPEN_ADMINISTRATIVELY_PROHIBITED, OPEN_CONNECT_FAILED

from core.database import async_session_factory
from gateway import server as gateway_server
Expand Down Expand Up @@ -95,6 +98,75 @@ def fake_provider(monkeypatch):
return FakeProcess


class FakeSandboxd(asyncssh.SSHServer):
"""Accepts any user and forwards each TCP channel on this machine."""

def begin_auth(self, username: str) -> bool:
return False

def connection_requested(self, dest_host, dest_port, orig_host, orig_port) -> bool:
return True


@pytest.fixture
async def sandbox_provider(monkeypatch):
sandboxd = await asyncssh.listen(
"127.0.0.1",
0,
server_host_keys=[asyncssh.generate_private_key("ssh-ed25519")],
server_factory=FakeSandboxd,
)
provider = SimpleNamespace(gateway_process_class=FakeProcess, tunnels=[])

async def open_gateway_tunnel(name):
tunnel = await asyncssh.connect(
"127.0.0.1",
sandboxd.get_port(),
username=name,
known_hosts=None,
config=None,
client_keys=None,
agent_path=None,
)
provider.tunnels.append(tunnel)
return tunnel

provider.open_gateway_tunnel = open_gateway_tunnel
monkeypatch.setattr(gateway_server, "get_vm_provider", lambda name: provider)
yield provider
sandboxd.close()


@pytest.fixture
async def forwarding_caller(gateway_settings, sandbox_provider):
caller_key = asyncssh.generate_private_key("ssh-ed25519")
await _insert_active_host("sb-forward", caller_key.export_public_key().decode())
server = await gateway_server.start(gateway_settings)
connection = await asyncssh.connect(
"127.0.0.1",
server.get_port(),
username="sb-forward",
client_keys=[caller_key],
known_hosts=None,
)
yield connection
connection.close()
server.close()


@pytest.fixture
async def loopback_service():
# The digest comes after the caller's EOF, so an answer proves half-close.
async def answer(reader, writer):
writer.write(hashlib.sha256(await reader.read()).digest())
await writer.drain()
writer.close()

service = await asyncio.start_server(answer, "127.0.0.1", 0)
yield service.sockets[0].getsockname()[1]
service.close()


@pytest.fixture
def local_provider(monkeypatch):
from gateway.tests.localprocess import LocalProcess
Expand Down Expand Up @@ -311,3 +383,47 @@ async def test_gateway_streams_binary_stdin_to_an_exec_and_delivers_the_status(

assert result.exit_status == 4
assert target.read_bytes() == payload


async def test_gateway_forwards_a_channel_to_sandbox_loopback(forwarding_caller, loopback_service):
payload = os.urandom(1 << 20)
reader, writer = await forwarding_caller.open_connection("127.0.0.1", loopback_service)
writer.write(payload)
writer.write_eof()

assert await reader.read() == hashlib.sha256(payload).digest()


@pytest.mark.parametrize("destination", ["10.0.0.1", "::1", "example.com"])
async def test_gateway_refuses_a_destination_outside_sandbox_loopback(
forwarding_caller, sandbox_provider, destination
):
with pytest.raises(asyncssh.ChannelOpenError) as refusal:
await forwarding_caller.open_connection(destination, 80)

assert refusal.value.code == OPEN_ADMINISTRATIVELY_PROHIBITED
assert not sandbox_provider.tunnels


async def test_gateway_shares_one_tunnel_and_closes_it_with_the_caller(
forwarding_caller, sandbox_provider, loopback_service
):
for _ in range(2):
reader, writer = await forwarding_caller.open_connection("localhost", loopback_service)
writer.write_eof()
await reader.read()
forwarding_caller.close()

[tunnel] = sandbox_provider.tunnels
await asyncio.wait_for(tunnel.wait_closed(), 5)


async def test_gateway_reports_a_tunnel_that_fails_to_open(forwarding_caller, sandbox_provider):
async def refuse(name):
raise ProviderTransportError("sandboxd is not running")

sandbox_provider.open_gateway_tunnel = refuse
with pytest.raises(asyncssh.ChannelOpenError) as refusal:
await forwarding_caller.open_connection("127.0.0.1", 80)

assert refusal.value.code == OPEN_CONNECT_FAILED
7 changes: 0 additions & 7 deletions src/gateway/tests/test_sftp.py
Original file line number Diff line number Diff line change
Expand Up @@ -147,13 +147,6 @@ async def read_one(i: int) -> bytes:
assert LocalProcess.open_count == 1


async def test_port_forwarding_stays_refused(connected):
# A direct-tcpip channel asks the gateway to open an outbound connection.
# The gateway serves no forwarding, thus the request is refused.
with pytest.raises(asyncssh.ChannelOpenError):
await connected.open_connection("127.0.0.1", 9)


async def test_a_command_runs_with_home_at_the_per_host_home(connected, tmp_path):
result = await connected.run('printf %s "$HOME"')
assert result.stdout == str(tmp_path / "sb-sftp")
Expand Down
8 changes: 8 additions & 0 deletions src/providers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,10 @@
from dataclasses import dataclass
from typing import ClassVar, NamedTuple, Self

import asyncssh

from providers.capabilities import ProxyInjection, SecretInjectionCapability
from providers.exceptions import ProviderTransportError


@dataclass(frozen=True)
Expand Down Expand Up @@ -104,6 +107,11 @@ async def create_vm(
@abc.abstractmethod
async def delete_vm(self, name: str) -> None: ...

async def open_gateway_tunnel(self, name: str) -> asyncssh.SSHClientConnection:
"""Open an SSH tunnel into a gateway host's sandbox. The gateway
forwards TCP channels through it."""
raise ProviderTransportError(f"{self.name} hosts have no gateway tunnel")

@abc.abstractmethod
async def diagnose(self) -> str:
"""One cheap read-only probe. Raise on failure: /doctor classifies the error."""
Expand Down
Loading
Loading