diff --git a/docs/deploy.md b/docs/deploy.md index a89a11f..d9ee9cf 100644 --- a/docs/deploy.md +++ b/docs/deploy.md @@ -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 -L 43123:127.0.0.1:43123 @ +``` + +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/`. The gateway makes this directory and moves into it first. Thus a caller writes files there — for example a `.gitconfig` or diff --git a/src/gateway/server.py b/src/gateway/server.py index 0527a05..6c77af3 100644 --- a/src/gateway/server.py +++ b/src/gateway/server.py @@ -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 @@ -19,21 +22,12 @@ _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): @@ -41,21 +35,26 @@ class GatewayConnection(asyncssh.SSHServer): 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 @@ -63,6 +62,36 @@ def connection_lost(self, exc: Exception | None) -> None: 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 @@ -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 @@ -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, @@ -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", diff --git a/src/gateway/tests/test_server.py b/src/gateway/tests/test_server.py index a6836cf..a81beca 100644 --- a/src/gateway/tests/test_server.py +++ b/src/gateway/tests/test_server.py @@ -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 @@ -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 @@ -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 diff --git a/src/gateway/tests/test_sftp.py b/src/gateway/tests/test_sftp.py index 4220a80..19f5bd1 100644 --- a/src/gateway/tests/test_sftp.py +++ b/src/gateway/tests/test_sftp.py @@ -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") diff --git a/src/providers/base.py b/src/providers/base.py index 3c46f43..a6fc2d0 100644 --- a/src/providers/base.py +++ b/src/providers/base.py @@ -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) @@ -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.""" diff --git a/src/providers/docker_sbx/provider.py b/src/providers/docker_sbx/provider.py index b426f4e..4af9551 100644 --- a/src/providers/docker_sbx/provider.py +++ b/src/providers/docker_sbx/provider.py @@ -4,6 +4,8 @@ from pathlib import Path from typing import ClassVar, Self +import asyncssh + from providers import environment from providers.base import VMCreateResult, VMProvider from providers.capabilities import TemplateCapability @@ -154,6 +156,23 @@ async def delete_vm(self, name: str) -> None: self._remove_sandbox_files(name) + async def open_gateway_tunnel(self, name: str) -> asyncssh.SSHClientConnection: + # sandboxd authenticates the OS user on its local socket, and sbx itself + # trusts the host key on first use. No key crosses a network. + try: + return await asyncssh.connect( + f"{name}.sbx", + username=self.settings.ssh_username, + proxy_command=["env", "SBX_NO_TELEMETRY=1", "sbx", "ssh", "proxy", f"{name}.sbx"], + known_hosts=None, + config=None, + client_keys=None, + agent_path=None, + preferred_auth="none", + ) + except (OSError, asyncssh.Error) as exc: + raise ProviderTransportError(f"sbx could not open a tunnel: {exc}") from exc + async def build_template_image( self, *,