Skip to content

proxy

Flight SSH proxy functionality.

FlightProxy

FlightProxy(url, backend_url, proxy_map, **kwargs)

Bases: FlightServerBase

Transparent Flight proxy that rewrites endpoint URIs based on mapping.

Rewrites FlightInfo endpoint location URIs on GetFlightInfo responses; all other read RPC methods are forwarded to the backend unchanged. This is useful for proxying a remote Flight server through SSH port forwards.

Write RPCs (DoPut/DoExchange) are intentionally not supported, so publishing through the proxy is not possible.

Parameters:

Name Type Description Default
url str

URL to listen for queries.

required
backend_url str

URL of remote info server to proxy.

required
proxy_map dict[str, str]

Dictionary mapping of endpoint locations to be replaced. It should be keyed on OLD_HOST:PORT with value as NEW_HOST:PORT. The URL scheme will be preserved.

required
Source code in arrakis/proxy.py
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
def __init__(
    self,
    url: str,
    backend_url: str,
    proxy_map: dict[str, str],
    **kwargs,
):
    super().__init__(url, **kwargs)
    self.proxy_map = proxy_map
    self.client = flight.connect(backend_url)
    logger.debug("flight proxy server initialized: %s -> %s", url, backend_url)
    logger.debug("proxy_map: %s", self.proxy_map)

SSHConnection

SSHConnection(destination)

Manage a background SSH connection

Parameters:

Name Type Description Default
destination str

SSH destination (see ssh(1) for form).

required
Source code in arrakis/proxy.py
189
190
191
192
193
194
195
196
197
def __init__(self, destination: str):
    self.destination = destination
    self.ctrl_dir = pathlib.Path(tempfile.mkdtemp())
    # destination may be a plain "[user@]host" instead of an
    # ssh:// URL, in which case netloc is empty
    ctrl_name = urlparse(self.destination).netloc or self.destination
    self.ctrl_path = self.ctrl_dir / ctrl_name
    self.ctrl = str(self.ctrl_path)
    self._forward_map: dict[str, str] = {}

forward_map property

forward_map

Dictionary of port forwards

REMOTE_HOSTPORT: LOCAL_HOSTPORT

close

close()

Close the connection

Source code in arrakis/proxy.py
292
293
294
295
296
def close(self):
    """Close the connection"""
    self.exec(["-O", "exit"], check=False, capture_output=True)
    # the control socket may not be removed immediately on exit
    shutil.rmtree(self.ctrl_dir, ignore_errors=True)

connect

connect()

connect to the ssh destination

Source code in arrakis/proxy.py
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
def connect(self):
    """connect to the ssh destination"""
    cmd = [
        "ssh",
        "-S",
        self.ctrl,
        "-M",
        "-o",
        "ControlPersist=yes",
        "-f",
        # "-N",
        self.destination,
        "sleep",
        "60",
    ]
    logger.debug(" ".join(cmd))
    subprocess.run(cmd, check=True)  # noqa S603

exec

exec(ssh_cmd=None, shell_cmd=None, **kwargs)

exec ssh command on control master

Source code in arrakis/proxy.py
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
def exec(
    self,
    ssh_cmd: list[str] | None = None,
    shell_cmd: list[str] | None = None,
    **kwargs,
):
    """exec ssh command on control master"""
    cmd = [
        "ssh",
        "-S",
        self.ctrl,
        "-o",
        "ControlMaster=no",
    ]
    if ssh_cmd:
        cmd += ssh_cmd
    cmd += [self.destination]
    if shell_cmd:
        cmd += shell_cmd
    logger.debug(" ".join(cmd))
    return subprocess.run(  # noqa S603
        cmd,
        **kwargs,
    )

forward_port

forward_port(remote_hostport, local_port=None, *, wait=False, timeout=10)

initiate a local port forward to remote location

Source code in arrakis/proxy.py
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
def forward_port(
    self,
    remote_hostport: str,
    local_port: int | None = None,
    *,
    wait: bool = False,
    timeout: float = 10,
):
    """initiate a local port forward to remote location"""
    if not local_port:
        # hacky way to find an unused port. the acquired port
        # could be stolen between when the socket releases it and
        # ssh tries to take it.
        s = socket.socket()
        s.bind(("localhost", 0))
        _, local_port = s.getsockname()
        s.close()

    local_hostport = f"localhost:{local_port}"

    forward = f"{local_hostport}:{remote_hostport}"
    logger.debug("forwarding port: %s", forward)
    self.exec(["-O", "forward", "-L", forward], check=True)

    if wait:
        # wait for forward to be established
        deadline = time.monotonic() + timeout
        while True:
            try:
                s = socket.create_connection(("localhost", local_port))
                s.close()
                break
            except OSError:
                if time.monotonic() > deadline:
                    msg = f"port forward not established: {forward}"
                    raise TimeoutError(msg) from None
                time.sleep(0.01)

    self.forward_map[remote_hostport] = local_hostport
    return local_hostport

SSHConnectionLike

Bases: Protocol

Interface required of connections used by ssh_proxy.

ssh_proxy

ssh_proxy(ssh_dest, arrakis_server=None, bind_address=None, connection_factory=SSHConnection)

Create Flight proxy server over SSH.

This is done by:

  1. Make SSH connection to the remote host.
  2. Determine initial Flight info server URL on the remote side.
  3. Retrieve all known endpoints from the remote info server.
  4. Set up local ssh port forwards to all endpoints.
  5. Launch Flight proxy server that rewrites endpoints to point to the local port forwards.

This should always be used as a context manager so that all connections are closed when done.

Note that publishing is not supported through the proxy; see FlightProxy.

Parameters:

Name Type Description Default
ssh_dest str

Remote SSH destination (see ssh(1) for form).

required
arrakis_server str | None

Remote Arrakis info server URL. If not specified it will be determined from the ARRAKIS_SERVER env var on the remote side.

None
bind_address str

Flight proxy server HOST:PORT. Defaults to "localhost:0" (random open port chosen on localhost).

None
connection_factory Callable[[str], SSHConnectionLike]

Factory producing the SSH connection for the given destination. Intended for substituting a fake connection in tests.

SSHConnection
Source code in arrakis/proxy.py
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
@contextlib.contextmanager
def ssh_proxy(
    ssh_dest: str,
    arrakis_server: str | None = None,
    bind_address: str | None = None,
    connection_factory: Callable[[str], SSHConnectionLike] = SSHConnection,
):
    """Create Flight proxy server over SSH.

    This is done by:

    0. Make SSH connection to the remote host.
    1. Determine initial Flight info server URL on the remote side.
    2. Retrieve all known endpoints from the remote info server.
    3. Set up local ssh port forwards to all endpoints.
    4. Launch Flight proxy server that rewrites endpoints to point to
       the local port forwards.

    This should always be used as a context manager so that all
    connections are closed when done.

    Note that publishing is not supported through the proxy; see
    `FlightProxy`.

    Parameters
    ----------
    ssh_dest : str
        Remote SSH destination (see ssh(1) for form).
    arrakis_server : str | None
        Remote Arrakis info server URL. If not specified it will be
        determined from the ARRAKIS_SERVER env var on the remote side.
    bind_address : str
        Flight proxy server HOST:PORT. Defaults to "localhost:0"
        (random open port chosen on localhost).
    connection_factory : Callable[[str], SSHConnectionLike]
        Factory producing the SSH connection for the given
        destination.  Intended for substituting a fake connection in
        tests.

    """
    logger.info("creating ssh proxy via %s", ssh_dest)

    # check/hold the requested server port
    bind_socket = socket.socket()
    if bind_address is None:
        bind_address = "localhost:0"
    bind_host, bind_port = bind_address.split(":")
    try:
        bind_socket.bind((bind_host, int(bind_port)))
    except OSError:
        msg = f"local address already in use: {bind_address}"
        raise OSError(msg) from None
    _, bind_port = bind_socket.getsockname()

    # initiate the ssh connection to the host as a context manager, so
    # that the connection is properly shut down if there are any
    # errors during setup.  closing() guarantees the held port is
    # released even if setup fails before the proxy server starts.
    with contextlib.closing(bind_socket), connection_factory(ssh_dest) as ssh:
        # if not specified, try to determine the remote server location
        # from the ARRAKIS_SERVER env var on the remote host
        if arrakis_server is None:
            logger.debug("resolving remote ARRAKIS_SERVER...")
            arrakis_server = (
                ssh.exec(
                    shell_cmd=["printenv", "ARRAKIS_SERVER"],
                    capture_output=True,
                    check=True,
                )
                .stdout.decode()
                .strip()
            )
            if not arrakis_server:
                msg = "Could not determine remote ARRAKIS_SERVER."
                raise ValueError(msg)

        # create an initial forward to the info server so that we can
        # query for the endpoint information
        backend_hostport = ssh.forward_port(parse_arrakis_url(arrakis_server).netloc)

        logger.debug("retrieving endpoints...")
        backend_url = f"grpc://{backend_hostport}"
        endpoints = Client(backend_url).endpoints()

        # setup forwards for all endpoints
        logger.debug("creating endpoint port forwarding...")
        for endpoint in endpoints:
            remote_hostport = urlparse(endpoint).netloc
            ssh.forward_port(remote_hostport)

        # start the proxy server
        logger.debug("starting Flight proxy server...")
        proxy_url = f"grpc://{bind_host}:{bind_port}"
        bind_socket.close()
        server = FlightProxy(proxy_url, backend_url, ssh.forward_map)

        executor = ThreadPoolExecutor(max_workers=1)
        future = executor.submit(server.serve)

        # confirm the server is actually serving, and surface any
        # startup error instead of yielding a dead proxy
        try:
            future.result(timeout=0.1)
        except FuturesTimeoutError:
            pass
        else:
            msg = f"flight proxy server failed to start: {proxy_url}"
            raise RuntimeError(msg)

        # yield context manager
        try:
            yield proxy_url
        finally:
            logger.info("shutting down ssh proxy...")
            server.shutdown()
            executor.shutdown()
            logger.info("done.")