Developer SDK

Python examples

Install the SDK · View source on GitHub ↗

socket/test_http.py

"""Local adapter integration tests; no account, host DNS, or public server."""
from contextlib import contextmanager
import gzip
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
import socket
import ssl
import subprocess
import tempfile
import threading
import time
from types import SimpleNamespace
import unittest
import httpcore
import httpx
import requests
from ur_http import UrHttpxTransport, UrRequestsAdapter, UrNetworkStream


class Handler(BaseHTTPRequestHandler):
    protocol_version = "HTTP/1.1"
    def log_message(self, *_):
        pass
    def do_GET(self):
        if self.path == "/redirect":
            self.send_response(302)
            self.send_header("Location", "/cookie")
            self.send_header("Set-Cookie", "example=yes; Path=/")
            data = b""
        else:
            self.send_response(200)
            data = self.headers.get("Cookie", "").encode() if self.path == "/cookie" else "héllo".encode()
        if self.path == "/gzip":
            data = gzip.compress(data)
            self.send_header("Content-Encoding", "gzip")
        self.send_header("Content-Type", "text/plain; charset=utf-8")
        self.send_header("Content-Length", str(len(data)))
        self.end_headers()
        self.wfile.write(data)
    def do_POST(self):
        data = self.rfile.read(int(self.headers.get("Content-Length", 0)))
        self.send_response(200)
        self.send_header("Content-Length", str(len(data)))
        self.end_headers()
        self.wfile.write(data)


class FixtureConn:
    # Exposes only the SDK methods. The adapters cannot access a file descriptor.
    def __init__(self, endpoint):
        self._socket = socket.create_connection(endpoint, timeout=5)
        self.closed = False
    def set_read_deadline(self, millis):
        self._socket.settimeout(None if not millis else max(.001, millis / 1000 - time.time()))
    set_write_deadline = set_read_deadline
    def read(self, size):
        data = self._socket.recv(size)
        return SimpleNamespace(data=data, eof=not data)
    def write(self, data):
        # Force partial native writes; the adapter must finish the byte sequence.
        return self._socket.send(data[:7])
    def close(self):
        self.closed = True
        self._socket.close()


class FixtureDevice:
    def __init__(self, endpoint):
        self.endpoint, self.names, self.conns = endpoint, [], []
    def dial(self, network, address, **_):
        assert network == "tcp"
        self.names.append(address)
        conn = FixtureConn(self.endpoint)
        self.conns.append(conn)
        return conn


@contextmanager
def server(tls=False):
    with tempfile.TemporaryDirectory() as temp:
        http = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
        cert = Path(temp) / "cert.pem"
        if tls:
            key = Path(temp) / "key.pem"
            subprocess.run(["openssl", "req", "-x509", "-newkey", "rsa:2048", "-nodes",
                            "-keyout", str(key), "-out", str(cert), "-days", "1",
                            "-subj", "/CN=socket.test", "-addext", "subjectAltName=DNS:socket.test"],
                           check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
            context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
            context.load_cert_chain(cert, key)
            http.socket = context.wrap_socket(http.socket, server_side=True)
        thread = threading.Thread(target=http.serve_forever, daemon=True)
        thread.start()
        try:
            yield FixtureDevice(http.server_address), ("https" if tls else "http") + "://socket.test:" + str(http.server_port), cert
        finally:
            http.shutdown()
            http.server_close()
            thread.join(5)


class AdapterTests(unittest.TestCase):
    def test_httpx_plain_and_verified_tls(self):
        for tls in (False, True):
            with self.subTest(tls=tls), server(tls) as (device, url, cert):
                context = ssl.create_default_context(cafile=str(cert)) if tls else None
                with httpx.Client(transport=UrHttpxTransport(device, context), trust_env=False, timeout=5) as client:
                    self.assertEqual(client.get(url).text, "héllo")
                    self.assertEqual(client.post(url, content="data").text, "data")
                    self.assertEqual(client.get(url + "/gzip").text, "héllo")
                self.assertTrue(all(c.closed for c in device.conns))
                self.assertTrue(all(name.startswith("socket.test:") for name in device.names))

    def test_requests_plain_and_verified_tls_redirect_cookie_and_compression(self):
        for tls in (False, True):
            with self.subTest(tls=tls), server(tls) as (device, url, cert):
                with requests.Session() as client:
                    client.trust_env = False
                    client.verify = str(cert) if tls else True
                    adapter = UrRequestsAdapter(device)
                    client.mount("http://", adapter)
                    client.mount("https://", adapter)
                    self.assertEqual(client.get(url, timeout=5).text, "héllo")
                    self.assertEqual(client.post(url, data="data", timeout=5).text, "data")
                    self.assertEqual(client.get(url + "/gzip", timeout=5).text, "héllo")
                    self.assertIn("example=yes", client.get(url + "/redirect", timeout=5).text)
                self.assertTrue(all(c.closed for c in device.conns))

    def test_untrusted_certificate_is_rejected(self):
        with server(True) as (device, url, _):
            with httpx.Client(transport=UrHttpxTransport(device), trust_env=False, timeout=5) as client:
                with self.assertRaises(httpx.NetworkError):
                    client.get(url)
            self.assertTrue(all(c.closed for c in device.conns))

    def test_partial_read_error_is_not_silently_lost(self):
        class Partial(FixtureConn):
            def __init__(self):
                pass
            def set_read_deadline(self, _):
                pass
            def read(self, _):
                error = OSError("read failed")
                error.data = b"prefix"
                raise error
        stream = UrNetworkStream(Partial())
        self.assertEqual(stream.read(20), b"prefix")
        with self.assertRaises(httpcore.ReadError):
            stream.read(20)


if __name__ == "__main__":
    unittest.main()