#!/usr/bin/env python3
"""Minimal SOCKS5 relay for the loopback lab: no authentication, CONNECT only.

It prints one line per request (address type and destination), so a run shows whether a client
reached the SOCKS layer at all. Bind it to 127.0.0.1 only; it is not meant for production use.

    python3 socks5_relay.py 1080
"""
import select
import socket
import socketserver
import struct
import sys

ATYP = {1: "IPv4", 3: "DOMAIN", 4: "IPv6"}


class Handler(socketserver.BaseRequestHandler):
    def read(self, size):
        data = b""
        while len(data) < size:
            chunk = self.request.recv(size - len(data))
            if not chunk:
                raise ConnectionError("client closed the connection")
            data += chunk
        return data

    def handle(self):
        try:
            _version, method_count = self.read(2)
            self.read(method_count)
            self.request.sendall(b"\x05\x00")  # no authentication
            _version, command, _reserved, atyp = self.read(4)
            if atyp == 1:
                host = socket.inet_ntoa(self.read(4))
            elif atyp == 3:
                host = self.read(self.read(1)[0]).decode("ascii", "replace")
            else:
                host = socket.inet_ntop(socket.AF_INET6, self.read(16))
            port = struct.unpack("!H", self.read(2))[0]
            print(f"CONNECT atyp={ATYP.get(atyp, atyp)} dest={host}:{port}", flush=True)
            if command != 1:
                self.request.sendall(b"\x05\x07\x00\x01" + bytes(6))  # command not supported
                return
            try:
                upstream = socket.create_connection((host, port), timeout=10)
            except OSError:
                self.request.sendall(b"\x05\x05\x00\x01" + bytes(6))  # connection refused
                return
            self.request.sendall(b"\x05\x00\x00\x01" + bytes(6))
            with upstream:
                pair = {self.request: upstream, upstream: self.request}
                while True:
                    readable, _, _ = select.select(list(pair), [], [], 60)
                    if not readable:
                        return
                    for source in readable:
                        data = source.recv(65536)
                        if not data:
                            return
                        pair[source].sendall(data)
        except (ConnectionError, OSError):
            return


class Server(socketserver.ThreadingTCPServer):
    allow_reuse_address = True
    daemon_threads = True


if __name__ == "__main__":
    port = int(sys.argv[1]) if len(sys.argv) > 1 else 1080
    with Server(("127.0.0.1", port), Handler) as server:
        print(f"SOCKS5 relay listening on 127.0.0.1:{port}", flush=True)
        server.serve_forever()
