diff --git a/doc/changelog.rst b/doc/changelog.rst index fd8ef94334..a737f37e25 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -20,14 +20,18 @@ PyMongo 4.19 brings a number of changes including: interpreter remain daemon threads and shutdown behavior is unchanged. Note that because these threads are non-daemon, a subinterpreter may block on teardown until any in-flight monitor work completes. -- Added the ``kms_connect_callback`` option on +- Added support for routing Key Management Service (KMS) requests for + Client-Side Field Level Encryption and Queryable Encryption through an HTTP + proxy, using the new ``kms_connect_callback`` option on :class:`~pymongo.encryption_options.AutoEncryptionOpts`, :class:`~pymongo.encryption.ClientEncryption`, and - :class:`~pymongo.asynchronous.encryption.AsyncClientEncryption`: a callable - that opens the connection to a Key Management Service (KMS) host for - Client-Side Field Level Encryption and Queryable Encryption, e.g. to route - KMS requests through a proxy. The driver performs the KMS TLS handshake over - the callback's connection, so verification still targets the KMS host. + :class:`~pymongo.asynchronous.encryption.AsyncClientEncryption`. The callback + opens the connection and the driver performs the KMS TLS handshake over it, so + verification still targets the KMS host rather than the proxy. For an ordinary + HTTP proxy, pass :class:`~pymongo.encryption_options.HTTPProxyKMSConnect` or + :class:`~pymongo.encryption_options.AsyncHTTPProxyKMSConnect` instead of + writing a callback. Both helpers accept custom ``CONNECT`` headers, e.g. + ``Proxy-Authorization`` for proxies that require authentication. - Added the ``srv_host_validator`` keyword argument to :class:`~pymongo.synchronous.mongo_client.MongoClient` and :class:`~pymongo.asynchronous.mongo_client.AsyncMongoClient`, an alternative to diff --git a/pymongo/_kms_connect_shared.py b/pymongo/_kms_connect_shared.py index c160657da9..ec2cd439e6 100644 --- a/pymongo/_kms_connect_shared.py +++ b/pymongo/_kms_connect_shared.py @@ -20,11 +20,22 @@ from __future__ import annotations +import asyncio +import base64 import contextlib +import functools +import re import socket -from collections.abc import Awaitable +import ssl +import threading +import time +import urllib.parse +from collections.abc import Awaitable, Mapping from dataclasses import dataclass -from typing import Any, Callable +from typing import Any, Callable, Optional + +from pymongo.errors import ConfigurationError +from pymongo.pool_shared import _close_late_socket @dataclass(frozen=True) @@ -35,6 +46,9 @@ class KMSConnectContext: :class:`socket.socket`. The driver performs the KMS TLS handshake over it, verifying against ``host`` rather than the peer actually reached. + Prefer :class:`HTTPProxyKMSConnect` or :class:`AsyncHTTPProxyKMSConnect` + over writing a callback. + :param host: Hostname of the KMS server, and the TLS verification target. :param port: Port of the KMS server. :param timeout: Seconds allowed for the connection: the default KMS @@ -70,6 +84,335 @@ def _close_rejected_kms_socket(obj: Any) -> None: close() -# Sphinx documents this class under pymongo.encryption_options, so the -# definition must claim that module name. +# Cap the CONNECT response header so a silent proxy cannot grow the buffer without bound. +_MAX_CONNECT_HEADER = 8192 + +# An RFC 7230 token: the grammar for a header field name. +_TOKEN_RE = re.compile(r"^[!#$%&'*+\-.^_`|~0-9A-Za-z]+$") + + +def _remaining(deadline: float) -> float: + """Seconds left before ``deadline``.""" + left = deadline - time.monotonic() + if left <= 0: + raise socket.timeout("timed out connecting through the proxy") + return left + + +class HTTPProxyKMSConnect: + """Route KMS connections through an HTTP proxy, for the synchronous API. + + Pass an instance as ``kms_connect_callback`` to reach KMS hosts through a + forward proxy that speaks HTTP ``CONNECT``:: + + from pymongo.encryption_options import AutoEncryptionOpts, HTTPProxyKMSConnect + + opts = AutoEncryptionOpts( + kms_providers={"aws": aws_creds}, + key_vault_namespace="keyvault.datakeys", + kms_connect_callback=HTTPProxyKMSConnect("http://proxy.example.com:8080"), + ) + + To reach the proxy over TLS, use an ``https`` proxy URL. Pass an + :class:`ssl.SSLContext` to customize trust; the default context is used + otherwise. It applies only to the proxy connection; KMS TLS is still + negotiated end to end:: + + import ssl + + proxy_tls = ssl.create_default_context(cafile="proxy-ca.pem") + callback = HTTPProxyKMSConnect("https://proxy.example.com:8443", proxy_tls) + + Use :class:`AsyncHTTPProxyKMSConnect` with the asynchronous API. + + :param proxy_url: URL of the proxy, in the same form as a proxy configured + for :mod:`urllib.request`, e.g. ``http://proxy.example.com:8080``. The + scheme must be ``http`` or ``https``, and the port defaults to 80 and + 443, respectively. Userinfo, e.g. + ``http://user:pass@proxy.example.com:8080``, is sent as a + ``Proxy-Authorization`` basic auth header. When sending credentials, + use an ``https`` proxy URL so they are not sent in cleartext. + :param ssl_context: Optional :class:`ssl.SSLContext` for connecting to the + proxy over TLS. Defaults to ``None``, meaning the default context for + an ``https`` proxy URL, or a plain connection for ``http``. + :param headers: Optional mapping of extra ``CONNECT`` request headers, + e.g. ``{"X-Trace-Id": "abc123"}``. ``Host`` is always set from the KMS + address, and ``Proxy-Authorization`` from any proxy URL userinfo. + + .. versionadded:: 4.19 + """ + + def __init__( + self, + proxy_url: str, + ssl_context: Optional[ssl.SSLContext] = None, + headers: Optional[Mapping[str, str]] = None, + ): + if not isinstance(proxy_url, str): + raise TypeError(f"proxy_url must be a string, not {type(proxy_url)}") + try: + split = urllib.parse.urlsplit(proxy_url) + except ValueError as exc: + # Malformed URLs, e.g. an unmatched IPv6 bracket, raise here + # rather than surfacing lazily from the parsed parts below. + raise ConfigurationError(f"invalid proxy_url: {proxy_url!r}") from exc + if split.scheme not in ("http", "https"): + raise ConfigurationError( + f"proxy_url must have an http or https scheme, not {proxy_url!r}" + ) + if not split.hostname: + raise ConfigurationError(f"proxy_url must include a host, not {proxy_url!r}") + if split.path not in ("", "/") or split.query or split.fragment: + raise ConfigurationError( + f"proxy_url must not include a path, query, or fragment: {proxy_url!r}" + ) + try: + port = split.port + except ValueError as exc: + raise ConfigurationError(f"invalid proxy_url: {proxy_url!r}") from exc + self.host = split.hostname + if port is None: + port = 443 if split.scheme == "https" else 80 + elif port == 0: + # An explicit port zero is never a valid proxy destination. + raise ConfigurationError(f"invalid proxy_url: {proxy_url!r}") + self.port = port + if split.scheme == "https": + self.ssl_context: Optional[ssl.SSLContext] = ( + ssl.create_default_context() if ssl_context is None else ssl_context + ) + else: + if ssl_context is not None: + raise ConfigurationError("ssl_context requires an https proxy_url") + self.ssl_context = None + self.headers = dict(headers) if headers else {} + for name, value in self.headers.items(): + if not isinstance(name, str) or not isinstance(value, str): + raise TypeError("proxy header names and values must be strings") + # Header fields become CONNECT request lines; a name outside the + # RFC 7230 token grammar, or CR/LF in a value, would corrupt or + # inject request lines. Use fullmatch: a match() of an anchored + # pattern accepts a trailing newline ($ matches just before it). + if not _TOKEN_RE.fullmatch(name): + raise ConfigurationError(f"invalid proxy header name: {name!r}") + if name.lower() == "host": + raise ConfigurationError("the Host CONNECT header is set from the KMS address") + if "\r" in value or "\n" in value: + raise ConfigurationError(f"invalid proxy header value for {name!r}") + if split.username is not None: + if any(name.lower() == "proxy-authorization" for name in self.headers): + raise ConfigurationError( + "proxy_url must not include userinfo when headers includes Proxy-Authorization" + ) + password = urllib.parse.unquote(split.password) if split.password else "" + creds = f"{urllib.parse.unquote(split.username)}:{password}".encode() + token = base64.b64encode(creds).decode("ascii") + self.headers["Proxy-Authorization"] = f"Basic {token}" + + def _tunnel(self, sock: socket.socket, context: KMSConnectContext, deadline: float) -> None: + # An IPv6 literal needs brackets to be a valid HTTP authority. + host = f"[{context.host}]" if ":" in context.host else context.host + target = f"{host}:{context.port}" + lines = [f"CONNECT {target} HTTP/1.1", f"Host: {target}"] + lines.extend(f"{name}: {value}" for name, value in self.headers.items()) + sock.sendall(("\r\n".join(lines) + "\r\n\r\n").encode()) + # Read a byte at a time: a bulk read could consume tunneled bytes from + # this same socket. Reapply the budget before each read so a trickling + # proxy cannot outlive the deadline. + response = bytearray() + while not response.endswith(b"\r\n\r\n"): + sock.settimeout(_remaining(deadline)) + chunk = sock.recv(1) + if not chunk: + raise OSError(f"proxy closed the connection while tunneling to {target}") + response += chunk + if len(response) > _MAX_CONNECT_HEADER: + raise OSError(f"proxy sent an oversized CONNECT response for {target}") + status = bytes(response).split(b"\r\n", 1)[0] + # A CONNECT is successful for any 2xx status, e.g. "HTTP/1.0 200" or + # "HTTP/1.1 201"; require a three-digit code and reject malformed lines. + parts = status.split(b" ", 2) + valid = ( + len(parts) >= 2 + and parts[0].startswith(b"HTTP/") + and len(parts[1]) == 3 + and parts[1].isdigit() + ) + if not valid or not 200 <= int(parts[1]) < 300: + raise OSError(f"proxy refused CONNECT to {target}: {status!r}") + + def _bridge(self, proxy: socket.socket) -> socket.socket: + """Relay a TLS proxy connection through a socketpair. + + Python cannot layer TLS over an :class:`ssl.SSLSocket`, so return the + plain end of a pair, using threads rather than tasks even in + :class:`AsyncHTTPProxyKMSConnect`: the event loop cannot read an + :class:`ssl.SSLSocket`. + """ + # Clear the CONNECT-phase timeout; the tunneled KMS request is governed + # by the driver's own timeout, not the elapsed connect budget. + proxy.settimeout(None) + driver_side, relay_side = socket.socketpair() + # Each socket is read by one relay and written by the other, and it + # closes only when both of those have ended: EOF propagates as a + # write-half shutdown, and the reverse direction may still carry + # a reply. + finished: dict[socket.socket, set[str]] = {relay_side: set(), proxy: set()} + phases_lock = threading.Lock() + + def relay(src: socket.socket, dst: socket.socket) -> None: + # Daemon threads: any error, including the ValueError an + # SSLSocket.shutdown can raise in the teardown race, ends the relay. + try: + while True: + buf = src.recv(16384) + if not buf: + break + dst.sendall(buf) + except (OSError, ValueError): + pass + finally: + with phases_lock: + finished[src].add("read") + finished[dst].add("write") + done = [s for s, marks in finished.items() if len(marks) == 2] + for s in done: + del finished[s] + for sock in done: + # Both directions over this socket have ended. + try: + sock.shutdown(socket.SHUT_RDWR) + except (OSError, ValueError): + pass + sock.close() + if dst not in done: + # Send EOF downstream without cutting the reverse + # direction, which may still deliver a reply. + try: + dst.shutdown(socket.SHUT_WR) + except (OSError, ValueError): + pass + + try: + for pair in ((relay_side, proxy), (proxy, relay_side)): + threading.Thread(target=relay, args=pair, daemon=True).start() + except BaseException: + # Unblock any thread that did start, then drop every socket. + # shutdown can raise the ValueError an SSLSocket shows in the + # teardown race; tolerate it here exactly as relay does. + for sock in (proxy, relay_side, driver_side): + try: + sock.shutdown(socket.SHUT_RDWR) + except (OSError, ValueError): + pass + sock.close() + raise + return driver_side + + def __call__( + self, context: KMSConnectContext, *, deadline: Optional[float] = None + ) -> socket.socket: + # A configurable KMS host could inject CR/LF into, or split the + # request line of, the CONNECT request. + if any(not c.isprintable() or c.isspace() for c in context.host): + raise ConfigurationError( + f"KMS host must not contain control characters or whitespace: {context.host!r}" + ) + # One deadline for all three phases; per-phase timeouts would multiply + # the caller's budget. + if deadline is None: + deadline = time.monotonic() + context.timeout + sock = self._connect_proxy(deadline) + try: + if self.ssl_context is not None: + sock.settimeout(_remaining(deadline)) + sock = self.ssl_context.wrap_socket(sock, server_hostname=self.host) + sock.settimeout(_remaining(deadline)) + self._tunnel(sock, context, deadline) + except BaseException: + sock.close() + raise + if self.ssl_context is None: + return sock + try: + return self._bridge(sock) + except BaseException: + sock.close() + raise + + def _connect_proxy(self, deadline: float) -> socket.socket: + # Recompute the budget per address, rather than socket.create_connection, + # which applies the timeout to every address. + # + # DNS resolution is not bounded by the deadline; a slow lookup can + # exceed the budget, as in the driver's own connect path. + last_error: Optional[OSError] = None + for family, socktype, proto, _, sockaddr in socket.getaddrinfo( + self.host, self.port, type=socket.SOCK_STREAM + ): + sock = socket.socket(family, socktype, proto) + try: + # Propagate the timeout from _remaining rather than report a + # connect error. + sock.settimeout(_remaining(deadline)) + except socket.timeout: + sock.close() + raise + try: + sock.connect(sockaddr) + except socket.timeout: + # Preserve the timeout type rather than report a generic + # connect error. + sock.close() + raise + except OSError as exc: + last_error = exc + sock.close() + continue + return sock + if last_error is None: + # getaddrinfo returned no usable addresses. + raise OSError(f"could not connect to proxy {self.host}:{self.port}") + raise OSError( + f"could not connect to proxy {self.host}:{self.port}: {last_error}" + ) from last_error + + +class AsyncHTTPProxyKMSConnect(HTTPProxyKMSConnect): + """Route KMS connections through an HTTP proxy, for the asynchronous API. + + Behaves exactly like :class:`HTTPProxyKMSConnect`, but is a coroutine + callable and runs the blocking connect in a thread so the event loop stays + free. + + .. versionadded:: 4.19 + """ + + async def __call__(self, context: KMSConnectContext) -> socket.socket: # type: ignore[override] + # Capture the deadline before scheduling so time spent queued behind + # other executor work counts against the KMS budget. + deadline = time.monotonic() + context.timeout + connect = functools.partial(super().__call__, context, deadline=deadline) + future = asyncio.get_running_loop().run_in_executor(None, connect) + try: + return await asyncio.wait_for(asyncio.shield(future), _remaining(deadline)) + except asyncio.CancelledError: + # The thread runs on regardless, so close the socket it returns. + future.add_done_callback(_close_late_socket) + raise + except TimeoutError: + if future.done(): + # The executor task timed out itself; report it directly. + raise + # The budget ran out with the thread still busy. The thread + # runs on regardless, so close the socket it returns, and + # report the deadline. + future.add_done_callback(_close_late_socket) + raise socket.timeout("timed out connecting through the proxy") from None + + +# Sphinx documents these classes under pymongo.encryption_options, the public +# import path, so the definitions must claim that module name. KMSConnectContext.__module__ = "pymongo.encryption_options" +HTTPProxyKMSConnect.__module__ = "pymongo.encryption_options" +AsyncHTTPProxyKMSConnect.__module__ = "pymongo.encryption_options" diff --git a/pymongo/asynchronous/_kms_connect.py b/pymongo/asynchronous/_kms_connect.py index 168d63425d..ab74919beb 100644 --- a/pymongo/asynchronous/_kms_connect.py +++ b/pymongo/asynchronous/_kms_connect.py @@ -113,7 +113,7 @@ async def _connect_kms( _close_rejected_kms_socket(sock) raise ConfigurationError( "kms_connect_callback must return a connected, unwrapped " - f"socket.socket, not {type(sock)}." + f"socket.socket, not {type(sock)}; consider AsyncHTTPProxyKMSConnect." ) # wrap_socket refuses a non-blocking socket, so normalize the mode here. try: diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 9c2d72604c..eca189a526 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -600,6 +600,10 @@ def __init__( the driver performs the KMS TLS handshake. Must be a coroutine function for the asynchronous API. The callback is responsible for honoring ``context.timeout``. + When a CSOT timeout is active, the driver stops waiting at the + deadline and closes any socket the callback yields later. For an + ordinary HTTP proxy, pass + :class:`~pymongo.encryption_options.AsyncHTTPProxyKMSConnect`. Defaults to ``None``, meaning the driver connects to KMS hosts directly. diff --git a/pymongo/encryption_options.py b/pymongo/encryption_options.py index 852ce37210..fdabecbbcb 100644 --- a/pymongo/encryption_options.py +++ b/pymongo/encryption_options.py @@ -36,7 +36,9 @@ _HAVE_PYMONGOCRYPT = False from bson import int64 from pymongo._kms_connect_shared import ( # noqa: F401 + AsyncHTTPProxyKMSConnect, AsyncKMSConnectCallback, + HTTPProxyKMSConnect, KMSConnectCallback, KMSConnectContext, ) diff --git a/pymongo/synchronous/_kms_connect.py b/pymongo/synchronous/_kms_connect.py index 12a8da93dc..78d8119557 100644 --- a/pymongo/synchronous/_kms_connect.py +++ b/pymongo/synchronous/_kms_connect.py @@ -113,7 +113,7 @@ def _connect_kms( _close_rejected_kms_socket(sock) raise ConfigurationError( "kms_connect_callback must return a connected, unwrapped " - f"socket.socket, not {type(sock)}." + f"socket.socket, not {type(sock)}; consider HTTPProxyKMSConnect." ) # wrap_socket refuses a non-blocking socket, so normalize the mode here. try: diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index d73aec4b68..1d9fae5888 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -597,6 +597,10 @@ def __init__( the driver performs the KMS TLS handshake. Must be a regular function. The callback is responsible for honoring ``context.timeout``. + When a CSOT timeout is active, the driver stops waiting at the + deadline and closes any socket the callback yields later. For an + ordinary HTTP proxy, pass + :class:`~pymongo.encryption_options.HTTPProxyKMSConnect`. Defaults to ``None``, meaning the driver connects to KMS hosts directly. diff --git a/test/asynchronous/test_encryption.py b/test/asynchronous/test_encryption.py index e22d185f77..bd07eb2c34 100644 --- a/test/asynchronous/test_encryption.py +++ b/test/asynchronous/test_encryption.py @@ -226,7 +226,8 @@ async def test_init_kms_tls_options(self): self.assertEqual(ctx.verify_mode, ssl.CERT_REQUIRED) -# KMS connect callback tests live in test_kms_connect.py. +# KMS connect callback unit tests live in test_kms_connect.py. The prose +# tests are in test_kms_connect_prose.py. class TestClientOptions(AsyncPyMongoTestCase): @@ -2007,7 +2008,8 @@ async def test_invalid_hostname_in_kms_certificate(self): await self.client_encrypted.create_data_key("aws", master_key=key) -# KMS connect callback tests live in test_kms_connect.py. +# KMS connect callback unit tests live in test_kms_connect.py. The prose +# tests are in test_kms_connect_prose.py. # https://github.com/mongodb/specifications/blob/master/source/client-side-encryption/tests/README.md#kms-tls-options-tests diff --git a/test/asynchronous/test_kms_connect_prose.py b/test/asynchronous/test_kms_connect_prose.py new file mode 100644 index 0000000000..44f2ac6942 --- /dev/null +++ b/test/asynchronous/test_kms_connect_prose.py @@ -0,0 +1,218 @@ +"""Integration tests for the KMS connect callback and HTTP proxy support. + +The unit tests live in ``test/test_kms_connect.py`` (a file that synchro +does not process). This module adds the integration tests, which run real +KMS traffic through a local proxy and are executed against both APIs via +the generated synchronous mirror. +""" + +from __future__ import annotations + +import asyncio +import http.client +import ssl +import unittest +from typing import Any + +import pytest + +from bson.binary import Binary +from pymongo.encryption_options import AsyncHTTPProxyKMSConnect, AutoEncryptionOpts +from pymongo.errors import EncryptionError +from test.asynchronous.test_encryption import OPTS, AsyncEncryptionIntegrationTest +from test.helpers_shared import AWS_CREDS, CA_PEM + +_IS_SYNC = False + +pytestmark = pytest.mark.encryption + +KMS_PROXY_HOST = "127.0.0.1" +KMS_PROXY_PORT = 9004 +KMS_TLS_PROXY_PORT = 9005 + +AWS_MASTER_KEY = { + "region": "us-east-1", + "key": "arn:aws:kms:us-east-1:579766882180:key/89fcc2c4-08b0-4bd9-9f25-e30687b580d0", +} + + +class TestKmsConnectCallbackProse(AsyncEncryptionIntegrationTest): + @unittest.skipUnless(any(AWS_CREDS.values()), "AWS environment credentials are not set") + async def asyncSetUp(self): + await super().asyncSetUp() + self.callback_calls: list[Any] = [] + + async def plain_callback(self, context): + self.callback_calls.append(context) + return await AsyncHTTPProxyKMSConnect(f"http://{KMS_PROXY_HOST}:{KMS_PROXY_PORT}")(context) + + def _proxy_tls_context(self): + ctx = ssl.create_default_context(cafile=CA_PEM) + ctx.check_hostname = False + # PYTHON-5040 tracks re-enabling verification once the test CA cert + # is fixed. The evergreen-tools CA lacks an Authority Key Identifier + # that newer OpenSSL requires, so verification fails on Windows 3.14. + ctx.verify_mode = ssl.CERT_NONE + return ctx + + async def tls_callback(self, context): + self.callback_calls.append(context) + callback = AsyncHTTPProxyKMSConnect( + f"https://{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", self._proxy_tls_context() + ) + return await callback(context) + + async def proxy_request(self, method, path, tls=False): + """Call the proxy's control endpoints and return the body.""" + if _IS_SYNC: + return self._proxy_request(method, path, tls) + return await asyncio.get_running_loop().run_in_executor( + None, self._proxy_request, method, path, tls + ) + + def _proxy_request(self, method, path, tls=False): + if tls: + conn = http.client.HTTPSConnection( + f"{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", context=self._proxy_tls_context() + ) + else: + conn = http.client.HTTPConnection(f"{KMS_PROXY_HOST}:{KMS_PROXY_PORT}") + try: + conn.request(method, path) + return conn.getresponse().read().decode() + finally: + conn.close() + + async def connect_count(self, tls=False): + body = await self.proxy_request("GET", "/metrics", tls=tls) + # One "key value" per line. The server also emits connect_target. + for line in body.splitlines(): + key, _, value = line.partition(" ") + if key == "connect_count": + return int(value) + raise AssertionError(f"no connect_count in metrics body: {body!r}") + + async def test_01_plain_http_proxy(self): + await self.proxy_request("POST", "/reset") + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(await self.connect_count(), 1) + + async def test_02_https_proxy(self): + await self.proxy_request("POST", "/reset", tls=True) + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.tls_callback, + ) + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(await self.connect_count(tls=True), 1) + + async def test_03_auto_encryption_through_proxy(self): + await self.client.keyvault.datakeys.drop() + await self.client.db.coll.drop() + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + data_key_id = await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + schema = { + "bsonType": "object", + "properties": { + "encrypted_string": { + "encrypt": { + "keyId": [data_key_id], + "bsonType": "string", + "algorithm": "AEAD_AES_256_CBC_HMAC_SHA_512-Deterministic", + } + } + }, + } + + await self.proxy_request("POST", "/reset") + opts = AutoEncryptionOpts( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + schema_map={"db.coll": schema}, + kms_connect_callback=self.plain_callback, + ) + client_encrypted = await self.async_rs_or_single_client(auto_encryption_opts=opts) + + await client_encrypted.db.coll.insert_one({"_id": 1, "encrypted_string": "hello"}) + decrypted = await client_encrypted.db.coll.find_one({"_id": 1}) + self.assertEqual(decrypted["encrypted_string"], "hello") + + raw = await self.client.db.coll.find_one({"_id": 1}) + self.assertIsInstance(raw["encrypted_string"], Binary) + + # The decrypt reuses the cached key, so exactly one KMS request follows + # the reset. + self.assertEqual(await self.connect_count(), 1) + + async def test_04_callback_error(self): + async def failing_callback(context): + raise OSError("proxy is on fire") + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=failing_callback, + ) + with self.assertRaisesRegex(EncryptionError, "proxy is on fire"): + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + + @unittest.skip( + "PYTHON-6037 ClientEncryption does not support timeoutMS, so the " + "callback always receives the default KMS connect timeout" + ) + async def test_05_callback_receives_timeout(self): + key_vault_client = await self.async_rs_or_single_client(timeoutMS=1000) + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + key_vault_client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + + self.assertTrue(self.callback_calls, "callback was never invoked") + for context in self.callback_calls: + # Checks only the spec's non-zero requirement, which cannot fail. + self.assertIsNotNone(context.timeout) + self.assertGreater(context.timeout, 0) + + async def test_06_retry_after_network_error(self): + state = {"calls": 0} + + async def flaky_callback(context): + state["calls"] += 1 + if state["calls"] == 1: + raise OSError("first attempt fails") + return await AsyncHTTPProxyKMSConnect(f"http://{KMS_PROXY_HOST}:{KMS_PROXY_PORT}")( + context + ) + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=flaky_callback, + ) + await encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(state["calls"], 2) diff --git a/test/test_encryption.py b/test/test_encryption.py index 2f9ed382ca..9ade8a8cae 100644 --- a/test/test_encryption.py +++ b/test/test_encryption.py @@ -226,7 +226,8 @@ def test_init_kms_tls_options(self): self.assertEqual(ctx.verify_mode, ssl.CERT_REQUIRED) -# KMS connect callback tests live in test_kms_connect.py. +# KMS connect callback unit tests live in test_kms_connect.py. The prose +# tests are in test_kms_connect_prose.py. class TestClientOptions(PyMongoTestCase): @@ -1999,7 +2000,8 @@ def test_invalid_hostname_in_kms_certificate(self): self.client_encrypted.create_data_key("aws", master_key=key) -# KMS connect callback tests live in test_kms_connect.py. +# KMS connect callback unit tests live in test_kms_connect.py. The prose +# tests are in test_kms_connect_prose.py. # https://github.com/mongodb/specifications/blob/master/source/client-side-encryption/tests/README.md#kms-tls-options-tests diff --git a/test/test_kms_connect.py b/test/test_kms_connect.py index 8209651776..4b530c0447 100644 --- a/test/test_kms_connect.py +++ b/test/test_kms_connect.py @@ -29,13 +29,16 @@ from __future__ import annotations import asyncio +import base64 import dataclasses import os import socket import ssl import threading +import time from asyncio.trsock import TransportSocket from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager from typing import Any from unittest import mock @@ -44,7 +47,13 @@ import pymongo from bson.codec_options import CodecOptions -from pymongo.encryption_options import _HAVE_PYMONGOCRYPT, AutoEncryptionOpts, KMSConnectContext +from pymongo.encryption_options import ( + _HAVE_PYMONGOCRYPT, + AsyncHTTPProxyKMSConnect, + AutoEncryptionOpts, + HTTPProxyKMSConnect, + KMSConnectContext, +) from pymongo.errors import ConfigurationError, ConnectionFailure, EncryptionError, NetworkTimeout from pymongo.pool_options import PoolOptions from pymongo.ssl_support import get_ssl_context @@ -98,6 +107,12 @@ async def connect(self, address, pool_options, callback, timeout): return await module._connect_kms(address, pool_options, callback, timeout) return module._connect_kms(address, pool_options, callback, timeout) + def proxy(self, proxy_url, tls_context=None, headers=None): + """The API's HTTP proxy KMS connect helper.""" + if self.is_async: + return AsyncHTTPProxyKMSConnect(proxy_url, tls_context, headers=headers) + return HTTPProxyKMSConnect(proxy_url, tls_context, headers=headers) + def callback(self, func): """Adapt a non-blocking ``func(context)`` to the API's callback form.""" if self.is_async: @@ -180,6 +195,31 @@ def _tls_server_context(cert=CLIENT_PEM): return ctx +def _insecure_client_context(): + # PYTHON-5040 tracks re-enabling verification: the evergreen-tools CA + # lacks an Authority Key Identifier newer OpenSSL requires. + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + return ctx + + +def _kms_context(host="kms.example.com", port=443, timeout=10): + """A KMSConnectContext with the defaults used throughout these tests.""" + return KMSConnectContext(host=host, port=port, timeout=timeout) + + +def _read_http_request(conn): + """Read until the blank line that ends a CONNECT request, or None on EOF.""" + request = b"" + while b"\r\n\r\n" not in request: + chunk = conn.recv(4096) + if not chunk: + return None + request += chunk + return request + + @contextmanager def _listen(backlog=1): listener = socket.socket() @@ -201,6 +241,72 @@ def _socketpair(): right.close() +@contextmanager +def _start_proxy(handler, backlog=1): + """Serve each accepted connection with ``handler(conn)`` in a daemon thread.""" + with _listen(backlog) as listener: + + def serve(): + for _ in range(backlog): + try: + conn, _ = listener.accept() + except OSError: + return + try: + handler(conn) + except OSError: + pass + finally: + conn.close() + + threading.Thread(target=serve, daemon=True).start() + yield listener.getsockname() + + +@contextmanager +def _record_and_reply(accepted, reply): + """A proxy that records each CONNECT request, replies ``reply``, and closes.""" + + def handler(conn): + request = _read_http_request(conn) + if request is None: + return + accepted.append(request) + conn.sendall(reply) + + with _start_proxy(handler) as addr: + yield addr + + +@contextmanager +def _tls_echo_proxy(delay=0): + """A TLS CONNECT proxy that replies 200, then echoes one tunneled read.""" + server_ctx = _tls_server_context() + + def handler(conn): + tls = server_ctx.wrap_socket(conn, server_side=True) + request = _read_http_request(tls) + if request is None: + return + tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + # The tunneled peer speaks only after the client does, as a TLS + # server would. The ``delay`` lets the reply outlast the CONNECT deadline. + if delay: + time.sleep(delay) + tls.sendall(b"echo:" + tls.recv(64)) + tls.close() + + with _start_proxy(handler) as addr: + yield addr + + +async def _echo_over_tunnel(api, sock): + sock.settimeout(10) + await api.offload(sock.sendall, b"ping") + data = await api.offload(sock.recv, 64) + assert data == b"echo:ping" + + @both_apis async def test_init_kms_connect_callback(api): opts = AutoEncryptionOpts({}, "k.d") @@ -379,6 +485,439 @@ async def test_cancelled_tls_wrap_closes_late_socket(api): assert left.fileno() == -1 +@both_apis +async def test_http_proxy_helper_tunnels_and_reports_refusal(api): + # Covers the CONNECT handshake without KMS credentials. + accepted: list[bytes] = [] + context = _kms_context() + + with _record_and_reply(accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n") as ( + host, + port, + ): + sock = await api.maybe_await(api.proxy(f"http://{host}:{port}")(context)) + with sock: + assert isinstance(sock, socket.socket) + assert accepted[0].split(b"\r\n")[0] == b"CONNECT kms.example.com:443 HTTP/1.1" + + with _record_and_reply(accepted, b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n") as ( + host, + port, + ): + with pytest.raises(OSError, match="refused CONNECT"): + await api.maybe_await(api.proxy(f"http://{host}:{port}")(context)) + + # Any 2xx status is a successful tunnel, not just HTTP/1.1 200. + with _record_and_reply(accepted, b"HTTP/1.0 200 Connection Established\r\n\r\n") as ( + host, + port, + ): + sock = await api.maybe_await(api.proxy(f"http://{host}:{port}")(context)) + with sock: + assert isinstance(sock, socket.socket) + + # A status code must be exactly three digits, with no zero padding. + for reply in (b"HTTP/1.1 2000 Evil\r\n\r\n", b"HTTP/1.1 00200 Evil\r\n\r\n"): + with _record_and_reply(accepted, reply) as (host, port): + with pytest.raises(OSError, match="refused CONNECT"): + await api.maybe_await(api.proxy(f"http://{host}:{port}")(context)) + + +@both_apis +async def test_control_characters_in_kms_host_are_rejected(api): + # Reject CR/LF in the configurable host before it reaches CONNECT. + callback = api.proxy("http://proxy.example.com:8080") + context = _kms_context(host="kms.example.com\r\nX-Injected: 1") + with pytest.raises(ConfigurationError, match="control characters or whitespace"): + await api.maybe_await(callback(context)) + # Whitespace would split the request line into extra tokens. + context = _kms_context(host="kms.example.com ") + with pytest.raises(ConfigurationError, match="control characters or whitespace"): + await api.maybe_await(callback(context)) + + +@both_apis +async def test_http_proxy_helper_sends_custom_headers(api): + # Extra CONNECT headers reach the proxy verbatim. + accepted: list[bytes] = [] + headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Trace-Id": "abc123"} + with _record_and_reply(accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n") as ( + host, + port, + ): + sock = await api.maybe_await( + api.proxy(f"http://{host}:{port}", headers=headers)(_kms_context()) + ) + with sock: + request = accepted[0] + assert request.split(b"\r\n")[0] == b"CONNECT kms.example.com:443 HTTP/1.1" + assert b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n" in request + assert b"\r\nX-Trace-Id: abc123\r\n" in request + assert request.count(b"\r\nHost: ") == 1 + + +@both_apis +async def test_http_proxy_helper_authenticates_to_the_proxy(api): + # The motivating case: 407 without credentials, 200 with them. + def handler(conn): + request = _read_http_request(conn) + if request is None: + return + if b"\r\nProxy-Authorization: Basic dXNlcjpwYXNz\r\n" in request: + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + else: + conn.sendall(b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n") + + context = _kms_context() + with _start_proxy(handler, backlog=2) as (host, port): + with pytest.raises(OSError, match="refused CONNECT"): + await api.maybe_await(api.proxy(f"http://{host}:{port}")(context)) + headers = {"Proxy-Authorization": "Basic dXNlcjpwYXNz"} + sock = await api.maybe_await(api.proxy(f"http://{host}:{port}", headers=headers)(context)) + with sock: + assert isinstance(sock, socket.socket) + + +@both_apis +async def test_http_proxy_helper_rejects_bad_headers(api): + for headers in [ + {"Bad\r\nName": "x"}, + {"Bad Name": "x"}, + {"Bad\tName": "x"}, + # A legal token followed by a newline: re's $ can match just before + # a trailing newline, so validation must require a full match. + {"X-Ok\n": "x"}, + {"X-Ok": "ok\r\nInjected: 1"}, + {"Host": "evil.example.com"}, + {"host": "evil.example.com"}, + {"": "x"}, + {"Bad:Name": "x"}, + ]: + with pytest.raises(ConfigurationError, match=r"proxy header|Host CONNECT header"): + api.proxy("http://proxy.example.com:8080", headers=headers) + + for headers in [{1: "x"}, {"X-Ok": 1}, {None: "x"}, {"X-Ok": None}]: + with pytest.raises(TypeError, match="must be strings"): + api.proxy("http://proxy.example.com:8080", headers=headers) + + +@both_apis +async def test_http_proxy_helper_accepts_legal_header_values(api): + # Colons and spaces are legal in values (e.g. auth schemes). Only + # CR/LF would let a value inject a request line. + callback = api.proxy( + "http://proxy.example.com:8080", + headers={"Proxy-Authorization": "Basic dXNlcjpwYXNz", "X-Token": "a: b"}, + ) + assert callback.headers == { + "Proxy-Authorization": "Basic dXNlcjpwYXNz", + "X-Token": "a: b", + } + + +@both_apis +async def test_proxy_url_is_parsed(api): + # The proxy URL has the same form as a proxy configured for urllib: + # scheme, optional userinfo, host, and optional port with defaults. + for url, host, port, tls in [ + ("http://proxy.example.com:8080", "proxy.example.com", 8080, False), + ("http://proxy.example.com", "proxy.example.com", 80, False), + ("https://proxy.example.com", "proxy.example.com", 443, True), + ]: + callback = api.proxy(url) + assert callback.host == host + assert callback.port == port + if tls: + assert callback.ssl_context is not None + else: + assert callback.ssl_context is None + + # An IPv6 literal is unwrapped for connecting. + callback = api.proxy("http://[::1]:8080") + assert callback.host == "::1" + assert callback.port == 8080 + + +@both_apis +async def test_https_proxy_url_defaults_to_the_default_context(api): + # An https proxy URL implies TLS, like urllib, using the default + # context unless one is passed. + assert api.proxy("https://proxy.example.com").ssl_context is not None + ctx = _insecure_client_context() + assert api.proxy("https://proxy.example.com", ctx).ssl_context is ctx + + +@both_apis +async def test_proxy_url_userinfo_authenticates_to_the_proxy(api): + # Userinfo becomes a Proxy-Authorization basic auth header, + # percent-decoded like urllib decodes it. + callback = api.proxy("http://user:p%40ss@proxy.example.com:8080") + expected = base64.b64encode(b"user:p@ss").decode() + assert callback.headers == {"Proxy-Authorization": f"Basic {expected}"} + + +@both_apis +async def test_proxy_url_is_validated(api): + for url in [ + "proxy.example.com:8080", # Missing scheme. + "ftp://proxy.example.com", # Not an HTTP(S) proxy. + "http://", # Missing host. + "http://proxy.example.com/path", + "http://proxy.example.com?x=1", + "http://proxy.example.com#frag", + "http://proxy.example.com:notaport", + "http://proxy.example.com:99999", + # An explicit port zero is never a valid proxy destination; it must + # not fall back to the scheme's default port. + "http://proxy.example.com:0", + "https://proxy.example.com:0", + "http://[::1", # Unmatched IPv6 bracket. + "http://[example.com]", # Invalid bracketed host. + ]: + with pytest.raises(ConfigurationError, match="proxy_url"): + api.proxy(url) + + # A TLS context is only meaningful for an https proxy URL. + with pytest.raises(ConfigurationError, match="https"): + api.proxy("http://proxy.example.com", _insecure_client_context()) + + # Userinfo and an explicit Proxy-Authorization header conflict. + with pytest.raises(ConfigurationError, match="Proxy-Authorization"): + api.proxy( + "http://user:pass@proxy.example.com", + headers={"Proxy-Authorization": "Basic dXNlcjpwYXNz"}, + ) + # Header validation runs before the userinfo handling, so a non-string + # name is a TypeError even when userinfo would also add a header. + with pytest.raises(TypeError, match="must be strings"): + api.proxy("http://user:pass@proxy.example.com", headers={1: "x"}) # type: ignore[dict-item] + with pytest.raises(TypeError, match="proxy_url"): + api.proxy(None) # type: ignore[arg-type] + + +@both_apis +async def test_tls_proxy_helper_bridges_the_tunnel(api): + # Covers the TLS-proxy path and the socketpair relay without KMS creds. + with _tls_echo_proxy() as (host, port): + sock = await api.maybe_await( + api.proxy(f"https://{host}:{port}", _insecure_client_context())(_kms_context()) + ) + with sock: + await _echo_over_tunnel(api, sock) + + +@both_apis +async def test_bridge_does_not_inherit_the_connect_deadline(api): + # The relay must outlast the much shorter CONNECT deadline. + with _tls_echo_proxy(delay=3.0) as (host, port): + sock = await api.maybe_await( + api.proxy(f"https://{host}:{port}", _insecure_client_context())( + _kms_context(timeout=2.0) + ) + ) + with sock: + await _echo_over_tunnel(api, sock) + + +@both_apis +async def test_bridge_half_close_does_not_lose_the_reply(api): + # A half-close is a write-side event: the relay propagates it as a + # write-half shutdown, and the reverse direction still delivers a reply. + listener = socket.create_server(("127.0.0.1", 0)) + proxy = socket.create_connection(listener.getsockname(), timeout=10) + peer, _ = listener.accept() + listener.close() + driver_side = HTTPProxyKMSConnect("http://proxy.example.com:8080")._bridge(proxy) + try: + driver_side.sendall(b"ping") + driver_side.shutdown(socket.SHUT_WR) + # The EOF reaches the tunnel peer, whose reply must still get through. + assert peer.recv(4096) == b"ping" + assert peer.recv(4096) == b"" + peer.sendall(b"echo:ping") + assert driver_side.recv(4096) == b"echo:ping" + finally: + driver_side.close() + peer.close() + + # Both relay directions have ended, so the relay closed the proxy socket. + for _ in range(50): + if proxy.fileno() == -1: + return + await asyncio.sleep(0.1) + assert proxy.fileno() == -1, "relay never closed the proxy socket" + + +@both_apis +async def test_proxy_closing_before_connect_reply_raises(api): + def handler(conn): + # Read the CONNECT request, then hang up without replying. + conn.recv(4096) + + with _start_proxy(handler) as (host, port): + with pytest.raises(OSError, match="proxy closed the connection"): + await api.maybe_await(api.proxy(f"http://{host}:{port}")(_kms_context())) + + +@async_only +async def test_cancelled_proxy_connect_closes_the_late_socket(api): + # A cancelled connect must close the socket the executor thread + # produces after the cancellation. + requested = threading.Event() + reply = threading.Event() + + def handler(conn): + conn.recv(4096) + requested.set() + if not reply.wait(10): + return + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + # Keep the connection open so the tunnel can complete its reads. + time.sleep(0.1) + + with _start_proxy(handler) as (host, port): + tunneled: list[socket.socket] = [] + original_tunnel = HTTPProxyKMSConnect._tunnel + + def spy_tunnel(self, sock, context, deadline): + tunneled.append(sock) + original_tunnel(self, sock, context, deadline) + + with mock.patch.object(HTTPProxyKMSConnect, "_tunnel", spy_tunnel): + callback = api.proxy(f"http://{host}:{port}") + task = asyncio.create_task(callback(_kms_context())) # type: ignore[arg-type] + waited = await api.offload(requested.wait, 10) + assert waited, "proxy never received the CONNECT request" + task.cancel("no longer needed") + with pytest.raises(asyncio.CancelledError): + await task + # Let the stub reply, completing the executor's future late. + reply.set() + await asyncio.sleep(0.5) + + assert len(tunneled) == 1 + assert tunneled[0].fileno() == -1, "late socket was left open" + + +@async_only +async def test_async_proxy_connect_wait_is_bounded_by_the_deadline(api): + # Time spent queued behind other executor work counts against the KMS + # budget: the await gives up at the deadline instead of waiting out the + # stall, and the thread's late socket, if any, is closed. + release = threading.Event() + + def blocker(): + release.wait(4.0) + + loop = asyncio.get_running_loop() + executor = ThreadPoolExecutor(max_workers=1) + loop.set_default_executor(executor) + try: + loop.run_in_executor(None, blocker) + callback = api.proxy("http://127.0.0.1:9") + start = time.monotonic() + task = asyncio.create_task(callback(_kms_context(timeout=1.0))) + await asyncio.sleep(0.2) # Let the coroutine queue its connect. + assert not task.done(), "connect finished before the deadline" + with pytest.raises(socket.timeout, match="timed out connecting through the proxy"): + await task + elapsed = time.monotonic() - start + assert elapsed < 2.5, f"the wait outlived the deadline: {elapsed:.1f}s" + finally: + release.set() + executor.shutdown(wait=False) + + +@both_apis +async def test_connect_timeout_is_not_reclassified(api): + # A connect that times out keeps its socket.timeout type instead of + # being reported as a generic connect error. + def timeout_connect(self, address): + raise socket.timeout("timed out") + + with mock.patch.object(socket.socket, "connect", timeout_connect): + with pytest.raises(socket.timeout): + HTTPProxyKMSConnect("http://127.0.0.1:9999")._connect_proxy(time.monotonic() + 10) + + +@both_apis +async def test_tunnel_keeps_bytes_sent_with_the_connect_reply(api): + # A proxy may coalesce its 200 with tunneled bytes. Reading past the header would drop them. + def handler(conn): + conn.recv(4096) + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\nearly-bytes") + + with _start_proxy(handler) as (host, port): + sock = await api.maybe_await(api.proxy(f"http://{host}:{port}")(_kms_context())) + with sock: + sock.settimeout(10) + data = await api.offload(sock.recv, 64) + assert data == b"early-bytes" + + +@both_apis +async def test_ipv6_host_is_bracketed_in_connect(api): + accepted: list[bytes] = [] + with _record_and_reply(accepted, b"HTTP/1.1 200 Connection Established\r\n\r\n") as ( + host, + port, + ): + sock = await api.maybe_await(api.proxy(f"http://{host}:{port}")(_kms_context(host="::1"))) + with sock: + assert accepted[0].split(b"\r\n")[0] == b"CONNECT [::1]:443 HTTP/1.1" + + +@both_apis +async def test_oversized_connect_response_is_rejected(api): + def handler(conn): + conn.recv(4096) + # Never sends the terminator. + while True: + conn.sendall(b"x" * 1024) + + with _start_proxy(handler) as (host, port): + with pytest.raises(OSError, match="oversized CONNECT response"): + await api.maybe_await(api.proxy(f"http://{host}:{port}")(_kms_context())) + + +@both_apis +async def test_remaining_raises_once_the_deadline_passes(api): + from pymongo._kms_connect_shared import _remaining + + assert _remaining(time.monotonic() + 5) > 0 + with pytest.raises(socket.timeout): + _remaining(time.monotonic() - 1) + + +@both_apis +async def test_bridge_failure_closes_the_proxy_socket(api): + # A failure inside _bridge must not strand the connected proxy socket. + server_ctx = _tls_server_context() + + def handler(conn): + tls = server_ctx.wrap_socket(conn, server_side=True) + tls.recv(4096) + tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + tls.close() + + captured = [] + + def failing_bridge(self, proxy): + captured.append(proxy) + raise OSError("no file descriptors") + + with _start_proxy(handler) as (host, port): + context = _kms_context() + + with mock.patch.object(HTTPProxyKMSConnect, "_bridge", failing_bridge): + with pytest.raises(OSError, match="no file descriptors"): + await api.maybe_await( + api.proxy(f"https://{host}:{port}", _insecure_client_context())(context) + ) + + assert captured[0].fileno() == -1, "proxy socket was left open" + + @async_only async def test_non_coroutine_callback_is_rejected(api): # A plain def must be rejected before it blocks the event loop. diff --git a/test/test_kms_connect_prose.py b/test/test_kms_connect_prose.py new file mode 100644 index 0000000000..f034d0a082 --- /dev/null +++ b/test/test_kms_connect_prose.py @@ -0,0 +1,216 @@ +"""Integration tests for the KMS connect callback and HTTP proxy support. + +The unit tests live in ``test/test_kms_connect.py`` (a file that synchro +does not process). This module adds the integration tests, which run real +KMS traffic through a local proxy and are executed against both APIs via +the generated synchronous mirror. +""" + +from __future__ import annotations + +import asyncio +import http.client +import ssl +import unittest +from typing import Any + +import pytest + +from bson.binary import Binary +from pymongo.encryption_options import AutoEncryptionOpts, HTTPProxyKMSConnect +from pymongo.errors import EncryptionError +from test.helpers_shared import AWS_CREDS, CA_PEM +from test.test_encryption import OPTS, EncryptionIntegrationTest + +_IS_SYNC = True + +pytestmark = pytest.mark.encryption + +KMS_PROXY_HOST = "127.0.0.1" +KMS_PROXY_PORT = 9004 +KMS_TLS_PROXY_PORT = 9005 + +AWS_MASTER_KEY = { + "region": "us-east-1", + "key": "arn:aws:kms:us-east-1:579766882180:key/89fcc2c4-08b0-4bd9-9f25-e30687b580d0", +} + + +class TestKmsConnectCallbackProse(EncryptionIntegrationTest): + @unittest.skipUnless(any(AWS_CREDS.values()), "AWS environment credentials are not set") + def setUp(self): + super().setUp() + self.callback_calls: list[Any] = [] + + def plain_callback(self, context): + self.callback_calls.append(context) + return HTTPProxyKMSConnect(f"http://{KMS_PROXY_HOST}:{KMS_PROXY_PORT}")(context) + + def _proxy_tls_context(self): + ctx = ssl.create_default_context(cafile=CA_PEM) + ctx.check_hostname = False + # PYTHON-5040 tracks re-enabling verification once the test CA cert + # is fixed. The evergreen-tools CA lacks an Authority Key Identifier + # that newer OpenSSL requires, so verification fails on Windows 3.14. + ctx.verify_mode = ssl.CERT_NONE + return ctx + + def tls_callback(self, context): + self.callback_calls.append(context) + callback = HTTPProxyKMSConnect( + f"https://{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", self._proxy_tls_context() + ) + return callback(context) + + def proxy_request(self, method, path, tls=False): + """Call the proxy's control endpoints and return the body.""" + if _IS_SYNC: + return self._proxy_request(method, path, tls) + return asyncio.get_running_loop().run_in_executor( + None, self._proxy_request, method, path, tls + ) + + def _proxy_request(self, method, path, tls=False): + if tls: + conn = http.client.HTTPSConnection( + f"{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", context=self._proxy_tls_context() + ) + else: + conn = http.client.HTTPConnection(f"{KMS_PROXY_HOST}:{KMS_PROXY_PORT}") + try: + conn.request(method, path) + return conn.getresponse().read().decode() + finally: + conn.close() + + def connect_count(self, tls=False): + body = self.proxy_request("GET", "/metrics", tls=tls) + # One "key value" per line. The server also emits connect_target. + for line in body.splitlines(): + key, _, value = line.partition(" ") + if key == "connect_count": + return int(value) + raise AssertionError(f"no connect_count in metrics body: {body!r}") + + def test_01_plain_http_proxy(self): + self.proxy_request("POST", "/reset") + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(self.connect_count(), 1) + + def test_02_https_proxy(self): + self.proxy_request("POST", "/reset", tls=True) + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.tls_callback, + ) + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(self.connect_count(tls=True), 1) + + def test_03_auto_encryption_through_proxy(self): + self.client.keyvault.datakeys.drop() + self.client.db.coll.drop() + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + data_key_id = encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + schema = { + "bsonType": "object", + "properties": { + "encrypted_string": { + "encrypt": { + "keyId": [data_key_id], + "bsonType": "string", + "algorithm": "AEAD_AES_256_CBC_HMAC_SHA_512-Deterministic", + } + } + }, + } + + self.proxy_request("POST", "/reset") + opts = AutoEncryptionOpts( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + schema_map={"db.coll": schema}, + kms_connect_callback=self.plain_callback, + ) + client_encrypted = self.rs_or_single_client(auto_encryption_opts=opts) + + client_encrypted.db.coll.insert_one({"_id": 1, "encrypted_string": "hello"}) + decrypted = client_encrypted.db.coll.find_one({"_id": 1}) + self.assertEqual(decrypted["encrypted_string"], "hello") + + raw = self.client.db.coll.find_one({"_id": 1}) + self.assertIsInstance(raw["encrypted_string"], Binary) + + # The decrypt reuses the cached key, so exactly one KMS request follows + # the reset. + self.assertEqual(self.connect_count(), 1) + + def test_04_callback_error(self): + def failing_callback(context): + raise OSError("proxy is on fire") + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=failing_callback, + ) + with self.assertRaisesRegex(EncryptionError, "proxy is on fire"): + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + + @unittest.skip( + "PYTHON-6037 ClientEncryption does not support timeoutMS, so the " + "callback always receives the default KMS connect timeout" + ) + def test_05_callback_receives_timeout(self): + key_vault_client = self.rs_or_single_client(timeoutMS=1000) + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + key_vault_client, + OPTS, + kms_connect_callback=self.plain_callback, + ) + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + + self.assertTrue(self.callback_calls, "callback was never invoked") + for context in self.callback_calls: + # Checks only the spec's non-zero requirement, which cannot fail. + self.assertIsNotNone(context.timeout) + self.assertGreater(context.timeout, 0) + + def test_06_retry_after_network_error(self): + state = {"calls": 0} + + def flaky_callback(context): + state["calls"] += 1 + if state["calls"] == 1: + raise OSError("first attempt fails") + return HTTPProxyKMSConnect(f"http://{KMS_PROXY_HOST}:{KMS_PROXY_PORT}")(context) + + encryption = self.create_client_encryption( + {"aws": AWS_CREDS}, + "keyvault.datakeys", + self.client, + OPTS, + kms_connect_callback=flaky_callback, + ) + encryption.create_data_key("aws", master_key=AWS_MASTER_KEY) + self.assertGreaterEqual(state["calls"], 2) diff --git a/tools/synchro.py b/tools/synchro.py index b5c6b2873a..4d2f865479 100644 --- a/tools/synchro.py +++ b/tools/synchro.py @@ -73,6 +73,7 @@ "AsyncClientEncryption": "ClientEncryption", "AsyncMongoCryptCallback": "MongoCryptCallback", "AsyncKMSConnectCallback": "KMSConnectCallback", + "AsyncHTTPProxyKMSConnect": "HTTPProxyKMSConnect", "AsyncExplicitEncrypter": "ExplicitEncrypter", "AsyncAutoEncrypter": "AutoEncrypter", "AsyncContextManager": "ContextManager",