init
This commit is contained in:
@@ -0,0 +1,137 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user