138 lines
4.5 KiB
Python
138 lines
4.5 KiB
Python
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)
|