diff --git a/resend/http_client_httpx.py b/resend/http_client_httpx.py index 32ea193..87df027 100644 --- a/resend/http_client_httpx.py +++ b/resend/http_client_httpx.py @@ -1,3 +1,4 @@ +import asyncio from typing import Any, Dict, List, Mapping, Optional, Tuple, Union import httpx @@ -8,10 +9,35 @@ class HTTPXClient(AsyncHTTPClient): """ Async HTTP client implementation using the httpx library. + + The client holds a single :class:`httpx.AsyncClient` so that the underlying + TCP connection (and its TLS handshake) is reused across requests. + + The underlying client is created lazily on first use rather than in + ``__init__``, because an ``httpx.AsyncClient`` binds its connection pool to + the running event loop and this class is instantiated at import time, when + no loop is running. If the running loop changes (for example, a second + ``asyncio.run(...)`` call), a fresh client is created for the new loop. + + Call :meth:`aclose` when the client is no longer needed. """ def __init__(self, timeout: int = 30): self._timeout = timeout + self._client: Optional[httpx.AsyncClient] = None + self._loop: Optional[asyncio.AbstractEventLoop] = None + + def _get_client(self) -> httpx.AsyncClient: + loop = asyncio.get_running_loop() + + if self._client is None or self._client.is_closed or self._loop is not loop: + # A client bound to a previous loop cannot be awaited on this one, + # and its pooled connections died with that loop, so it is dropped + # rather than closed here. + self._client = httpx.AsyncClient(timeout=self._timeout) + self._loop = loop + + return self._client async def request( self, @@ -22,26 +48,34 @@ async def request( files: Optional[Dict[str, Any]] = None, data: Optional[Dict[str, str]] = None, ) -> Tuple[bytes, int, Mapping[str, str]]: + client = self._get_client() + try: - async with httpx.AsyncClient(timeout=self._timeout) as client: - if files is not None: - resp = await client.request( - method=method, - url=url, - headers=headers, - files=files, - data=data, - ) - else: - resp = await client.request( - method=method, - url=url, - headers=headers, - json=json if data is None else None, - data=data, - ) - return resp.content, resp.status_code, resp.headers + if files is not None: + resp = await client.request( + method=method, + url=url, + headers=headers, + files=files, + data=data, + ) + else: + resp = await client.request( + method=method, + url=url, + headers=headers, + json=json if data is None else None, + data=data, + ) + return resp.content, resp.status_code, resp.headers except httpx.RequestError as e: # This gets caught by the async request.perform() method # and raises a ResendError with the error type "HttpClientError" raise RuntimeError(f"Request failed: {e}") from e + + async def aclose(self) -> None: + """Close the underlying client and release pooled connections.""" + if self._client is not None and not self._client.is_closed: + await self._client.aclose() + self._client = None + self._loop = None diff --git a/resend/http_client_requests.py b/resend/http_client_requests.py index 8caf308..83363fc 100644 --- a/resend/http_client_requests.py +++ b/resend/http_client_requests.py @@ -8,10 +8,16 @@ class RequestsClient(HTTPClient): """ This is the default HTTP client implementation using the requests library. + + The client holds a single :class:`requests.Session` so that the underlying + TCP connection (and its TLS handshake) is reused across requests. Call + :meth:`close` when the client is no longer needed, or use it as a context + manager. """ def __init__(self, timeout: int = 30): self._timeout = timeout + self._session = requests.Session() def request( self, @@ -24,7 +30,7 @@ def request( ) -> Tuple[bytes, int, Mapping[str, str]]: try: if files is not None: - resp = requests.request( + resp = self._session.request( method=method, url=url, headers=headers, @@ -33,7 +39,7 @@ def request( timeout=self._timeout, ) else: - resp = requests.request( + resp = self._session.request( method=method, url=url, headers=headers, @@ -46,3 +52,13 @@ def request( # This gets caught by the request.perform() method # and raises a ResendError with the error type "HttpClientError" raise RuntimeError(f"Request failed: {e}") from e + + def close(self) -> None: + """Close the underlying session and release pooled connections.""" + self._session.close() + + def __enter__(self) -> "RequestsClient": + return self + + def __exit__(self, *args: Any) -> None: + self.close() diff --git a/tests/http_client_connection_reuse_test.py b/tests/http_client_connection_reuse_test.py new file mode 100644 index 0000000..58cdc36 --- /dev/null +++ b/tests/http_client_connection_reuse_test.py @@ -0,0 +1,199 @@ +import asyncio +import json +import socketserver +import threading +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from http.server import BaseHTTPRequestHandler +from typing import Any, Iterator, List, Tuple + +import resend +from resend.http_client_httpx import HTTPXClient +from resend.http_client_requests import RequestsClient + + +class _ThreadedServer(socketserver.ThreadingTCPServer): + daemon_threads = True + allow_reuse_address = True + + +@contextmanager +def serve() -> Iterator[Tuple[str, List[int]]]: + """Run a local HTTP server and record how many TCP connections it accepts. + + Yields the base URL and the connection log. One entry is appended per + accepted connection, so reuse shows up as a shorter log than the number of + requests made. + """ + connections: List[int] = [] + lock = threading.Lock() + + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def setup(self) -> None: + # setup() runs once per accepted TCP connection, not per request. + with lock: + connections.append(1) + super().setup() + + def do_POST(self) -> None: + length = int(self.headers.get("Content-Length", 0)) + sent = json.loads(self.rfile.read(length)) + + # Echo the subject back so a response delivered to the wrong + # caller is detectable. + body = json.dumps({"id": sent["subject"]}).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, format: str, *args: Any) -> None: + pass + + httpd = _ThreadedServer(("127.0.0.1", 0), Handler) + threading.Thread(target=httpd.serve_forever, daemon=True).start() + + try: + yield f"http://127.0.0.1:{httpd.server_address[1]}", connections + finally: + httpd.shutdown() + httpd.server_close() + + +SEND_PARAMS: resend.Emails.SendParams = { + "from": "hello@example.com", + "to": ["world@example.com"], + "subject": "Hi!", + "html": "hi", +} + + +class TestSyncConnectionReuse: + def setup_method(self) -> None: + self._original_client = resend.default_http_client + self._original_url = resend.api_url + resend.api_key = "re_test" + + def teardown_method(self) -> None: + resend.default_http_client = self._original_client + resend.api_url = self._original_url + resend.api_key = None + + def test_reuses_a_single_connection(self) -> None: + with serve() as (url, connections): + resend.api_url = url + client = RequestsClient() + resend.default_http_client = client + + try: + for _ in range(5): + resend.Emails.send(SEND_PARAMS) + finally: + client.close() + + assert len(connections) == 1 + + def test_close_is_idempotent(self) -> None: + client = RequestsClient() + client.close() + client.close() + + def test_shared_session_keeps_responses_separate_across_threads(self) -> None: + """The session is shared, so each caller must still get its own response.""" + with serve() as (url, _connections): + resend.api_url = url + client = RequestsClient() + resend.default_http_client = client + + def send(index: int) -> Tuple[int, str]: + params: resend.Emails.SendParams = { + "from": "hello@example.com", + "to": ["world@example.com"], + "subject": f"msg-{index}", + "html": "hi", + } + return index, resend.Emails.send(params)["id"] + + try: + with ThreadPoolExecutor(max_workers=8) as pool: + results = list(pool.map(send, range(40))) + finally: + client.close() + + assert all(sent_id == f"msg-{index}" for index, sent_id in results) + + def test_works_as_a_context_manager(self) -> None: + with serve() as (url, connections): + resend.api_url = url + + with RequestsClient() as client: + resend.default_http_client = client + for _ in range(3): + resend.Emails.send(SEND_PARAMS) + + assert len(connections) == 1 + + +class TestAsyncConnectionReuse: + def setup_method(self) -> None: + self._original_client = resend.default_async_http_client + self._original_url = resend.api_url + resend.api_key = "re_test" + + def teardown_method(self) -> None: + resend.default_async_http_client = self._original_client + resend.api_url = self._original_url + resend.api_key = None + + def test_reuses_a_single_connection(self) -> None: + with serve() as (url, connections): + resend.api_url = url + client = HTTPXClient() + resend.default_async_http_client = client + + async def send_many() -> None: + try: + for _ in range(5): + await resend.Emails.send_async(SEND_PARAMS) + finally: + await client.aclose() + + asyncio.run(send_many()) + + assert len(connections) == 1 + + def test_recreates_the_client_when_the_event_loop_changes(self) -> None: + """A client cached from a closed loop must not be reused on a new one.""" + with serve() as (url, connections): + resend.api_url = url + client = HTTPXClient() + resend.default_async_http_client = client + + async def send_one() -> None: + await resend.Emails.send_async(SEND_PARAMS) + + # Two separate loops. The second must not fail on the dead pool. + asyncio.run(send_one()) + asyncio.run(send_one()) + + # One connection per loop, and no error raised. + assert len(connections) == 2 + + def test_aclose_allows_a_later_request(self) -> None: + with serve() as (url, connections): + resend.api_url = url + client = HTTPXClient() + resend.default_async_http_client = client + + async def send_close_send() -> None: + await resend.Emails.send_async(SEND_PARAMS) + await client.aclose() + await resend.Emails.send_async(SEND_PARAMS) + await client.aclose() + + asyncio.run(send_close_send()) + + assert len(connections) == 2 diff --git a/tests/request_test.py b/tests/request_test.py index 56e2b3d..2e772b0 100644 --- a/tests/request_test.py +++ b/tests/request_test.py @@ -11,7 +11,7 @@ class TestResendRequest(unittest.TestCase): - @patch("resend.http_client_requests.requests.request") + @patch("resend.http_client_requests.requests.Session.request") @patch("resend.api_key", new="test_key") def test_request_idempotency_key_is_set(self, mock_requests: MagicMock) -> None: mock_response = Mock() @@ -41,7 +41,7 @@ def test_request_idempotency_key_is_set(self, mock_requests: MagicMock) -> None: self.assertEqual(headers["User-Agent"], f"resend-python:{get_version()}") self.assertEqual(headers["Idempotency-Key"], "abc-123") - @patch("resend.http_client_requests.requests.request") + @patch("resend.http_client_requests.requests.Session.request") @patch("resend.api_key", new="test_key") def test_request_idempotency_key_is_not_set(self, mock_requests: MagicMock) -> None: mock_response = Mock() @@ -72,7 +72,7 @@ def test_request_idempotency_key_is_not_set(self, mock_requests: MagicMock) -> N "Idempotency-Key", headers, "Idempotency-Key should not be set" ) - @patch("resend.http_client_requests.requests.request") + @patch("resend.http_client_requests.requests.Session.request") @patch("resend.api_key", new="test_key") def test_non_json_preserves_http_status_when_client_error( self, mock_requests: MagicMock @@ -101,7 +101,7 @@ def test_non_json_preserves_http_status_when_client_error( self.assertIn("text/html", err.message) self.assertEqual(err.headers.get("retry-after"), "2") - @patch("resend.http_client_requests.requests.request") + @patch("resend.http_client_requests.requests.Session.request") @patch("resend.api_key", new="test_key") def test_non_json_preserves_http_status_when_server_error( self, mock_requests: MagicMock @@ -125,7 +125,7 @@ def test_non_json_preserves_http_status_when_server_error( self.assertEqual(err.code, 503) self.assertEqual(err.error_type, "application_error") - @patch("resend.http_client_requests.requests.request") + @patch("resend.http_client_requests.requests.Session.request") @patch("resend.api_key", new="test_key") def test_non_json_falls_back_to_500_when_status_is_success( self, mock_requests: MagicMock @@ -150,7 +150,7 @@ def test_non_json_falls_back_to_500_when_status_is_success( self.assertEqual(err.error_type, "application_error") self.assertIn("text/html", err.message) - @patch("resend.http_client_requests.requests.request") + @patch("resend.http_client_requests.requests.Session.request") @patch("resend.api_key", new="test_key") def test_invalid_json_preserves_http_status(self, mock_requests: MagicMock) -> None: mock_response = Mock() @@ -173,7 +173,7 @@ def test_invalid_json_preserves_http_status(self, mock_requests: MagicMock) -> N self.assertEqual(err.error_type, "application_error") self.assertEqual(err.message, "Failed to decode JSON response") - @patch("resend.http_client_requests.requests.request") + @patch("resend.http_client_requests.requests.Session.request") @patch("resend.api_key", new="test_key") def test_json_error_uses_http_status_when_body_omits_status_code( self, mock_requests: MagicMock