from __future__ import annotations import contextlib import functools import http.server import socket import threading from pathlib import Path from urllib.request import Request from n_m3u8dl_py.http import HttpClient, _Socks5HTTPSHandler def test_http_client_routes_requests_through_no_auth_socks5_proxy(tmp_path: Path) -> None: (tmp_path / "payload.txt").write_text("proxied", encoding="utf-8") with _http_server(tmp_path) as target_url, _socks5_proxy() as (proxy_url, address_types): target_url = target_url.replace("127.0.0.1", "localhost") data, resolved, _ = HttpClient(proxy=proxy_url, use_system_proxy=False).get_bytes(f"{target_url}/payload.txt") assert data == b"proxied" assert resolved.endswith("/payload.txt") assert address_types == [3] def test_socks5_https_handler_uses_the_standard_library_context_contract() -> None: handler = _Socks5HTTPSHandler("127.0.0.1", 1080) captured: dict[str, object] = {} def do_open(http_class, request, **kwargs): captured.update(kwargs) return "opened" handler.do_open = do_open # type: ignore[method-assign] assert handler.https_open(Request("https://example.test/video.m3u8")) == "opened" assert captured == {"context": handler._context} @contextlib.contextmanager def _http_server(directory: Path): handler = functools.partial(http.server.SimpleHTTPRequestHandler, directory=str(directory)) server = http.server.ThreadingHTTPServer(("127.0.0.1", 0), handler) thread = threading.Thread(target=server.serve_forever, daemon=True) thread.start() try: yield f"http://127.0.0.1:{server.server_port}" finally: server.shutdown() server.server_close() thread.join() @contextlib.contextmanager def _socks5_proxy(): listener = socket.socket(socket.AF_INET, socket.SOCK_STREAM) listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) listener.bind(("127.0.0.1", 0)) listener.listen() listener.settimeout(0.2) stop = threading.Event() workers: list[threading.Thread] = [] address_types: list[int] = [] def serve() -> None: while not stop.is_set(): try: client, _ = listener.accept() except TimeoutError: continue except OSError: if stop.is_set(): return raise worker = threading.Thread(target=_handle_client, args=(client, address_types), daemon=True) workers.append(worker) worker.start() thread = threading.Thread(target=serve, daemon=True) thread.start() try: yield f"socks5://127.0.0.1:{listener.getsockname()[1]}", address_types finally: stop.set() listener.close() thread.join() for worker in workers: worker.join() def _handle_client(client: socket.socket, address_types: list[int]) -> None: with client: assert _receive_exact(client, 2) == b"\x05\x01" assert _receive_exact(client, 1) == b"\x00" client.sendall(b"\x05\x00") header = _receive_exact(client, 4) assert header[:3] == b"\x05\x01\x00" address_types.append(header[3]) if header[3] == 1: host = socket.inet_ntoa(_receive_exact(client, 4)) elif header[3] == 3: host = _receive_exact(client, _receive_exact(client, 1)[0]).decode("idna") else: raise AssertionError("Unexpected SOCKS address type") port = int.from_bytes(_receive_exact(client, 2), "big") with socket.create_connection((host, port)) as upstream: client.sendall(b"\x05\x00\x00\x01\x00\x00\x00\x00\x00\x00") left = threading.Thread(target=_copy, args=(client, upstream), daemon=True) right = threading.Thread(target=_copy, args=(upstream, client), daemon=True) left.start() right.start() left.join() right.join() def _copy(source: socket.socket, target: socket.socket) -> None: try: while data := source.recv(64 * 1024): target.sendall(data) except OSError: pass finally: try: target.shutdown(socket.SHUT_WR) except OSError: pass def _receive_exact(connection: socket.socket, count: int) -> bytes: chunks: list[bytes] = [] while count: chunk = connection.recv(count) if not chunk: raise OSError("Socket closed") chunks.append(chunk) count -= len(chunk) return b"".join(chunks)