a73x

test/native_forward_remote.py

Ref:   Size: 19.0 KiB   History

#!/usr/bin/env python3
"""Opt-in real SSH and direct-QUIC forwarding against one private remote fixture.

Safety latch: this script contacts no host unless MUX_FORWARD_REMOTE_ENABLE=1.
The fixture owns every remote path, daemon socket, QUIC key, UDP port, service,
and copied binary; it never touches the remote user's mux state.
"""
import contextlib
import hashlib
import json
import os
from pathlib import Path
import secrets
import shlex
import shutil
import socket
import subprocess
import sys
import tempfile
import threading
import time

sys.dont_write_bytecode = True
from native_forward import (ForwardRig, PreResponseTransportStartupError, echo_roundtrip,
                            distinct_ports, http_get, listener_reserved)
from native_tiling import eventually, require


REMOTE_DEFAULT = 'ubuntu@192.168.0.107'
REMOTE_ADDRESS = '192.168.0.107'
REMOTE_PATH = '/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin'
REMOTE_PYTHON = '/usr/bin/python3'

REMOTE_SERVICE = r'''import hashlib,http.server,json,os,socket,socketserver,sys,tempfile,threading
root,out=sys.argv[1:]
token=hashlib.sha256(root.encode()).hexdigest()[:20]
def port():
 s=socket.socket();s.setsockopt(socket.SOL_SOCKET,socket.SO_REUSEADDR,1);s.bind(("127.0.0.1",0));n=s.getsockname()[1];s.close();return n
class H(http.server.BaseHTTPRequestHandler):
 def do_GET(self):
  b=("native-forward-remote-"+token+"\n").encode();self.send_response(200);self.send_header("Content-Length",str(len(b)));self.send_header("Connection","close");self.end_headers();self.wfile.write(b)
 def log_message(self,*a):pass
class E(socketserver.BaseRequestHandler):
 def handle(self):
  while True:
   b=self.request.recv(65536)
   if not b:self.request.shutdown(socket.SHUT_WR);return
   self.request.sendall(b)
servers=[]
for p,h in ((port(),H),(port(),E)):
 s=socketserver.ThreadingTCPServer(("127.0.0.1",p),h);s.daemon_threads=True;threading.Thread(target=s.serve_forever,daemon=True).start();servers.append(s)
fd,tmp=tempfile.mkstemp(prefix="ports-",dir=os.path.dirname(out));os.write(fd,json.dumps({"http":servers[0].server_address[1],"echo":servers[1].server_address[1],"token":token}).encode());os.fsync(fd);os.close(fd);os.replace(tmp,out)
threading.Event().wait()
'''


class RemoteFixture:
    def __init__(self, local_mux):
        self.local_mux = str(Path(local_mux).resolve())
        self.host = os.environ.get('MUX_FORWARD_REMOTE', REMOTE_DEFAULT)
        self.ssh, self.scp = shutil.which('ssh'), shutil.which('scp')
        require(self.ssh and self.scp, 'remote forwarding requires ssh and scp')
        self.local = Path(tempfile.mkdtemp(prefix='mux-forward-ssh-'))
        self.remote = None
        self.daemon_pid = None
        self.quic_port = None
        self.cleanup_errors = []
        self.marker = secrets.token_hex(24)
        self._rollback = contextlib.ExitStack()
        self._rollback.callback(self.close)
        try:
            self._make_local_key()
            self._create_remote_root()
            self.remote_mux = self.remote + '/bin/mux'
            self.remote_sock = self.remote + '/mux.sock'
            self.remote_key = self.remote + '/quic.key'
            self.remote_oracle = self.remote + '/os_oracle.sh'
            self.env_words = {
                'PATH': REMOTE_PATH, 'HOME': self.remote + '/home',
                'XDG_CONFIG_HOME': self.remote + '/xdg/config',
                'XDG_STATE_HOME': self.remote + '/xdg/state',
                'XDG_CACHE_HOME': self.remote + '/xdg/cache',
                'XDG_RUNTIME_DIR': self.remote + '/xdg/runtime', 'SHELL': '/bin/sh',
                'MUX_KEY_FILE': self.remote_key, 'MUX_SOCK': self.remote_sock,
            }
            self.run('mkdir -p ' + ' '.join(shlex.quote(self.remote + '/' + p) for p in
                     ('bin', 'home', 'xdg/config', 'xdg/state', 'xdg/cache', 'xdg/runtime')))
            for source, destination in ((self.local_mux, self.remote_mux),
                                        (str(self.local_key), self.remote_key),
                                        (str(Path('test/os_oracle.sh').resolve()), self.remote_oracle)):
                subprocess.run([self.scp, '-q', source, self.host + ':' + destination], check=True, timeout=30)
            self.run('chmod 700 ' + shlex.quote(self.remote_mux) + ' && chmod 600 ' +
                     shlex.quote(self.remote_key) + ' ' + shlex.quote(self.remote_oracle))
            service = self.local / 'remote-services.py'
            service.write_text(REMOTE_SERVICE)
            subprocess.run([self.scp, '-q', str(service), self.host + ':' + self.remote + '/services.py'], check=True, timeout=30)
            self.run_env(shlex.join([REMOTE_PYTHON, self.remote + '/services.py', self.remote, self.remote + '/ports.json']) +
                         ' > ' + shlex.quote(self.remote + '/services.log') + ' 2>&1 & printf %s "$!" > ' +
                         shlex.quote(self.remote + '/services.pid'))
            self.ports = eventually(self._read_ports, 'remote loopback services did not publish ports JSON', seconds=15)
            self.quic_port = self._high_udp_port()
            self.start_daemon()
            self._rollback.pop_all()
        except BaseException:
            self._rollback.close()
            raise

    def _make_local_key(self):
        key_env = os.environ.copy()
        for key in ('XDG_CONFIG_HOME', 'XDG_STATE_HOME', 'XDG_CACHE_HOME', 'XDG_RUNTIME_DIR', 'HOME'):
            path = self.local / key
            path.mkdir(mode=0o700, exist_ok=True)
            key_env[key] = str(path)
        subprocess.run([self.local_mux, 'd', 'keygen'], env=key_env, check=True, capture_output=True, text=True, timeout=15)
        self.local_key = self.local / 'XDG_CONFIG_HOME' / 'mux' / 'key'
        require(self.local_key.is_file() and self.local_key.stat().st_size > 0, 'fixture keygen did not create a private key')

    def _create_remote_root(self):
        result = self.run('umask 077; mktemp -d /tmp/mux-forward-%s-XXXXXXXX' % os.getuid())
        candidate = result.stdout.strip()
        require(candidate.startswith('/tmp/mux-forward-') and '\n' not in candidate, 'remote mktemp returned an unsafe fixture root')
        self.remote = candidate
        validate = ('import os,stat,sys; p,m=sys.argv[1:]; s=os.lstat(p); '
                    'assert stat.S_ISDIR(s.st_mode) and s.st_uid==os.getuid(); os.chmod(p,0o700); '
                    'fd=os.open(p+"/.mux-forward-owner",os.O_WRONLY|os.O_CREAT|os.O_EXCL,0o600); '
                    'os.write(fd,m.encode()); os.close(fd); print("OK")')
        result = self.run(shlex.join([REMOTE_PYTHON, '-c', validate, candidate, self.marker]), check=False)
        require(result.returncode == 0 and result.stdout.strip() == 'OK', 'remote fixture root failed owner/mode/marker validation')

    def env_prefix(self):
        return 'env -u MUX_SESSION ' + ' '.join(k + '=' + shlex.quote(v) for k, v in self.env_words.items())

    def run(self, command, check=True, timeout=30):
        return subprocess.run([self.ssh, self.host, command], text=True, capture_output=True, check=check, timeout=timeout)

    def run_env(self, command, check=True, timeout=30):
        return self.run(self.env_prefix() + ' /bin/sh -c ' + shlex.quote(command), check, timeout)

    def run_mux(self, verb, *args, stdout_log=None, check=True):
        command = shlex.join([self.remote_mux, 'd', verb, *args, '--sock', self.remote_sock])
        if stdout_log:
            command += ' > ' + shlex.quote(self.remote + '/' + stdout_log) + ' 2>&1'
        return self.run_env(command, check)

    def _high_udp_port(self):
        code = ('import random,socket; '\
                's=socket.socket(socket.AF_INET,socket.SOCK_DGRAM); '\
                's.bind(("%s",random.randrange(40000,60000))); print(s.getsockname()[1]); s.close()' % REMOTE_ADDRESS)
        result = self.run(shlex.join([REMOTE_PYTHON, '-c', code]))
        port = int(result.stdout.strip())
        require(40000 <= port < 60000, 'remote fixture did not reserve a high UDP port')
        return port

    def _udp_hex(self):
        packed = socket.inet_aton(REMOTE_ADDRESS)
        return packed[::-1].hex().upper() + ':' + format(self.quic_port, '04X')

    def _oracle(self, body, check=True):
        return self.run_env('. ' + shlex.quote(self.remote_oracle) + '; ' + body, check=check, timeout=15)

    def _record_daemon_pid(self):
        # The PID is accepted only when the OS oracle sees both the copied
        # executable and this concrete Unix listener in its fd table.
        body = ('want=$(real_path ' + shlex.quote(self.remote_mux) + ') || exit 41; n=0; found=; '
                'for d in /proc/[0-9]*; do p=${d##*/}; pid_alive "$p" || continue; '
                'exe=$(pid_exe "$p") || continue; [ "$exe" = "$want" ] || continue; '
                'pid_holds_unix_sock "$p" ' + shlex.quote(self.remote_sock) + ' || continue; '
                'n=$((n+1)); found=$p; done; [ "$n" -eq 1 ] || exit 42; printf "%s\\n" "$found"')
        result = self._oracle(body, check=False)
        require(result.returncode == 0 and result.stdout.strip().isdigit(),
                'OS oracle could not uniquely identify private daemon executable/socket owner: ' + result.stderr)
        self.daemon_pid = int(result.stdout.strip())
        print('OS-ORACLE: daemon pid=%d owns copied mux executable and %s' %
              (self.daemon_pid, self.remote_sock), flush=True)

    def _assert_daemon_live(self):
        self._record_daemon_pid()
        body = ('pid_alive ' + str(self.daemon_pid) + ' && pid_holds_unix_sock ' +
                str(self.daemon_pid) + ' ' + shlex.quote(self.remote_sock) + ' && udp_local_bound ' + self._udp_hex())
        result = self._oracle(body, check=False)
        require(result.returncode == 0, 'OS oracle did not see private daemon socket and QUIC UDP listener')

    def start_daemon(self):
        self.run_mux('start', '-d', '--quic', REMOTE_ADDRESS + ':' + str(self.quic_port), '--key', self.remote_key,
                     stdout_log='daemon.log')
        eventually(lambda: self._daemon_ready(), 'private remote QUIC daemon did not start', seconds=15)

    def _daemon_ready(self):
        try:
            self._assert_daemon_live()
            return True
        except AssertionError:
            return False

    def stop_daemon_verified(self):
        require(self.daemon_pid is not None, 'private daemon PID was never OS-verified')
        stop = self.run_mux('stop', check=False)
        body = ('! pid_alive ' + str(self.daemon_pid) + ' && [ ! -e ' + shlex.quote(self.remote_sock) +
                ' ] && ! udp_local_bound ' + self._udp_hex())
        result = self._oracle(body, check=False)
        require(stop.returncode == 0, 'private daemon stop command failed: ' + stop.stderr)
        require(result.returncode == 0, 'private daemon termination lacked OS evidence (pid/socket/UDP): ' + result.stderr)
        print('OS-ORACLE: daemon pid=%d gone; Unix socket removed; QUIC UDP listener released' %
              self.daemon_pid, flush=True)
        self.daemon_pid = None

    def _read_ports(self):
        result = self.run(shlex.join([REMOTE_PYTHON, '-c', 'import sys; print(open(sys.argv[1]).read())', self.remote + '/ports.json']), check=False)
        if result.returncode != 0:
            return None
        try:
            ports = json.loads(result.stdout)
            return ports if all(isinstance(ports.get(k), int) and 0 < ports[k] < 65536 for k in ('http', 'echo')) and isinstance(ports.get('token'), str) else None
        except json.JSONDecodeError:
            return None

    def _stop_owned_service(self):
        code = '''import os,sys,time
root,marker=sys.argv[1:]; pid=int(open(root+'/services.pid').read().strip())
cmd=open('/proc/%d/cmdline'%pid,'rb').read(); assert root.encode() in cmd and (root+'/services.py').encode() in cmd and open(root+'/.mux-forward-owner').read()==marker
os.kill(pid,15)
for _ in range(50):
 try: os.kill(pid,0)
 except ProcessLookupError: sys.exit(0)
 time.sleep(.1)
sys.exit(43)'''
        result = self.run(shlex.join([REMOTE_PYTHON, '-c', code, self.remote, self.marker]), check=False, timeout=10)
        require(result.returncode == 0, 'remote service identity/stop verification failed: ' + result.stderr)

    def _remove_owned_root(self):
        code = '''import os,shutil,stat,sys
root,marker=sys.argv[1:]; s=os.lstat(root)
assert stat.S_ISDIR(s.st_mode) and s.st_uid==os.getuid() and (s.st_mode&0o777)==0o700
assert open(root+'/.mux-forward-owner').read()==marker
assert not os.path.exists(root+'/mux.sock')
shutil.rmtree(root)'''
        result = self.run(shlex.join([REMOTE_PYTHON, '-c', code, self.remote, self.marker]), check=False, timeout=15)
        require(result.returncode == 0, 'remote root ownership validation refused cleanup: ' + result.stderr)

    def _copy_log(self, failure_logs, name):
        result = self.run('cat ' + shlex.quote(self.remote + '/' + name), check=False, timeout=10)
        (Path(failure_logs) / ('remote-' + name)).write_text(result.stdout + result.stderr)

    def close(self, failure_logs=None):
        if not self.remote:
            shutil.rmtree(self.local, ignore_errors=True)
            return
        if failure_logs is not None:
            for name in ('daemon.log', 'services.log'):
                try: self._copy_log(failure_logs, name)
                except BaseException as error: self.cleanup_errors.append('copy ' + name + ': ' + repr(error))
        daemon_stopped = self.daemon_pid is None
        if not daemon_stopped:
            try:
                self.stop_daemon_verified(); daemon_stopped = True
            except BaseException as error:
                self.cleanup_errors.append('stop/verify private daemon: ' + repr(error))
        try: self._stop_owned_service()
        except BaseException as error: self.cleanup_errors.append('stop verified private service: ' + repr(error))
        if daemon_stopped:
            try: self._remove_owned_root()
            except BaseException as error: self.cleanup_errors.append('remove verified private root: ' + repr(error))
        else:
            self.cleanup_errors.append('remote root retained because daemon termination was not OS-verified: ' + self.remote)
        shutil.rmtree(self.local, ignore_errors=True)
        if not self.cleanup_errors:
            print('PASS: remote fixture cleanup (daemon PID/socket/UDP released; services and root removed)', flush=True)


def exact_http(port, expected_body, deadline=None):
    code, body = http_get(port, deadline)
    require(code == 200 and body == expected_body,
            'forwarded HTTP response was not exact HTTP 200 fixture identity: ' +
            repr((code, body[:160])))


def wait_exact_http(port, body, message, seconds=15):
    deadline = time.monotonic() + seconds
    while time.monotonic() < deadline:
        try:
            exact_http(port, body, deadline)
            return
        except PreResponseTransportStartupError:
            time.sleep(min(.04, max(0, deadline - time.monotonic())))
    raise AssertionError(message)


def exact_http_api_regression():
    original = http_get
    try:
        globals()['http_get'] = lambda *_: (200, b'fixture')
        exact_http(1, b'fixture')
        for response, message in (((201, b'fixture'), 'non-200'),
                                  ((200, b'wrong-body'), 'wrong-body')):
            globals()['http_get'] = lambda *_, response=response: response
            try:
                exact_http(1, b'fixture')
            except AssertionError:
                pass
            else:
                raise AssertionError('exact_http accepted a ' + message + ' tuple response')
    finally:
        globals()['http_get'] = original


def run_mode(rig, remote, label, target_args, recover=False):
    local_http, local_echo = distinct_ports(2)
    rules = [(local_http, remote.ports['http']), (local_echo, remote.ports['echo'])]
    print('COMMAND: ' + shlex.join([rig.muxg, *target_args, '--forward', f'{local_http}:{remote.ports["http"]}', '--forward', f'{local_echo}:{remote.ports["echo"]}', '--session', 'forward']), flush=True)
    rig.launch_forward(target_args, rules, label + '-gui')
    body = ('native-forward-remote-' + remote.ports['token'] + '\n').encode()
    wait_exact_http(local_http, body, label + ' forwarding did not reach exact remote HTTP identity')
    payloads = [os.urandom(1024 * 1024 + 31 + i) for i in range(3)]
    replies, errors = [None] * len(payloads), []
    lock = threading.Lock()
    def worker(index):
        try: replies[index] = echo_roundtrip(local_echo, payloads[index], half_close=True)
        except BaseException as error:
            with lock: errors.append(error)
    workers = [threading.Thread(target=worker, args=(i,)) for i in range(3)]
    for worker in workers: worker.start()
    for worker in workers:
        worker.join(20)
        require(not worker.is_alive(), label + ' concurrent >1MiB stream hung')
    require(not errors, label + ' concurrent stream failure: ' + repr(errors))
    require(replies == payloads, label + ' concurrent >1MiB bidirectional half-closed streams corrupted bytes')
    require(listener_reserved(local_http), label + ' listener was not reserved')
    checksums = [hashlib.sha256(p).hexdigest() for p in payloads]
    rig.ok(label + ' exact HTTP 200 and three concurrent >1MiB half-close streams sha256=' + ','.join(checksums))
    if recover:
        remote.stop_daemon_verified()
        eventually(lambda: not _http_or_false(local_http), label + ' old listener did not become unavailable')
        require(listener_reserved(local_http), label + ' interruption released listener reservation')
        remote.start_daemon()
        wait_exact_http(local_http, body, label + ' did not recover after whole-daemon restart')
        rig.ok(label + ' whole-daemon restart recovery retains listener reservation and reconnects')
    rig.quit()


def _http_or_false(port):
    try: return bool(http_get(port))
    except PreResponseTransportStartupError: return False


def main():
    require(len(sys.argv) == 3, 'usage: native_forward_remote.py RELEASE_MUX RELEASE_MUXG')
    exact_http_api_regression()
    require(os.environ.get('MUX_FORWARD_REMOTE_ENABLE') == '1', 'set MUX_FORWARD_REMOTE_ENABLE=1 to contact the authorized remote fixture host')
    remote = rig = None
    try:
        remote = RemoteFixture(sys.argv[1])
        rig = ForwardRig(sys.argv[1], sys.argv[2])
        via = shlex.join([remote.ssh, remote.host, remote.remote_mux, 'd', 'proxy', '--sock', remote.remote_sock])
        run_mode(rig, remote, 'real SSH stdio', ['--via', via])
        run_mode(rig, remote, 'direct QUIC', ['quic://' + REMOTE_ADDRESS + ':' + str(remote.quic_port), '--key', str(remote.local_key)], recover=True)
        print('PASS: remote native forwarding', flush=True)
    except BaseException:
        if rig: rig.failure_artifacts()
        raise
    finally:
        failed = sys.exc_info()[0] is not None
        try:
            # GUI quit precedes final remote stop, preventing reconnect races.
            if rig: rig.close()
        finally:
            if remote:
                remote.close(rig.root if rig and failed else None)
                if remote.cleanup_errors:
                    raise RuntimeError('remote fixture cleanup errors: ' + '; '.join(remote.cleanup_errors))


if __name__ == '__main__':
    main()