diff --git a/doc/changelog.rst b/doc/changelog.rst index e752f9bcb2..c07212e3e3 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -51,6 +51,17 @@ PyMongo 4.18 brings a number of changes including: :meth:`~pymongo.synchronous.database.Database.aggregate`, and :meth:`~pymongo.asynchronous.collection.AsyncCollection.list_search_indexes` and :meth:`~pymongo.synchronous.collection.Collection.list_search_indexes`. +- 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`. 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. Changes in Version 4.17.0 (2026/04/20) -------------------------------------- diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 524ae45c11..cd0fbaf888 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -19,7 +19,9 @@ import asyncio import contextlib import enum +import inspect import socket +import ssl import time as time # noqa: PLC0414 # needed in sync version import uuid import weakref @@ -63,7 +65,9 @@ from pymongo.common import CONNECT_TIMEOUT from pymongo.daemon import _spawn_daemon from pymongo.encryption_options import ( + AsyncKMSConnectCallback, AutoEncryptionOpts, + KMSConnectContext, RangeOpts, TextOpts, check_min_pymongocrypt, @@ -82,6 +86,7 @@ from pymongo.pool_options import PoolOptions from pymongo.pool_shared import ( _async_configured_socket, + _async_wrap_socket_tls, _raise_connection_failure, ) from pymongo.read_concern import ReadConcern @@ -112,9 +117,65 @@ _KEY_VAULT_OPTS = CodecOptions(document_class=RawBSONDocument) -async def _connect_kms(address: _Address, opts: PoolOptions) -> Union[socket.socket, _sslConn]: +def _close_rejected_kms_socket(obj: Any) -> None: + """Close a rejected kms_connect_callback return value, if it can be closed. + + The caller may have handed us a live socket, and nothing else will close it: + _connect_kms raises before its result reaches the caller's ``finally``. The + value can be anything a user returned, so closing is strictly best effort. + """ + close = getattr(obj, "close", None) + if callable(close): + with contextlib.suppress(Exception): + close() + + +async def _connect_kms( + address: _Address, + opts: PoolOptions, + kms_connect_callback: Optional[AsyncKMSConnectCallback], + timeout: float, +) -> Union[socket.socket, _sslConn]: + if kms_connect_callback is None: + try: + return await _async_configured_socket(address, opts) + except Exception as exc: + _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) + + # TLS targets address, not the peer, so verification follows the KMS host. + result = kms_connect_callback( + KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) + ) + # The synchronous module takes a regular function and awaits nothing. + if not _IS_SYNC and not inspect.isawaitable(result): + _close_rejected_kms_socket(result) + raise ConfigurationError( + "kms_connect_callback must be a coroutine function for the async " + f"API, but returned {type(result)}." + ) + sock = await result + if not isinstance(sock, socket.socket) or isinstance(sock, ssl.SSLSocket): + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return a connected, unwrapped " + f"socket.socket, not {type(sock)}; consider HTTPProxyKMSConnect." + ) + # wrap_socket refuses a non-blocking socket, so normalize the mode here. try: - return await _async_configured_socket(address, opts) + sock.getpeername() + except OSError: + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return an already connected socket." + ) from None + if sock.getsockopt(socket.SOL_SOCKET, socket.SO_TYPE) != socket.SOCK_STREAM: + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return a stream socket, not a datagram one." + ) + sock.settimeout(opts.socket_timeout) + try: + return await _async_wrap_socket_tls(sock, address, opts) except Exception as exc: _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) @@ -184,20 +245,26 @@ async def kms_request(self, kms_context: MongoCryptKmsContext) -> None: False, # disable_ocsp_endpoint_check _IS_SYNC, ) - # CSOT: set timeout for socket creation. + address = parse_host(endpoint, _HTTPS_PORT) + sleep_u = kms_context.usleep + if sleep_u: + sleep_sec = float(sleep_u) / 1e6 + await asyncio.sleep(sleep_sec) + # CSOT: set timeout for socket creation. After the retry backoff above, + # so the budget reflects what the sleep consumed. connect_timeout = max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001) opts = PoolOptions( connect_timeout=connect_timeout, socket_timeout=connect_timeout, ssl_context=ctx, ) - address = parse_host(endpoint, _HTTPS_PORT) - sleep_u = kms_context.usleep - if sleep_u: - sleep_sec = float(sleep_u) / 1e6 - await asyncio.sleep(sleep_sec) try: - conn = await _connect_kms(address, opts) + conn = await _connect_kms( + address, + opts, + self.opts._kms_connect_callback, + connect_timeout, + ) try: await async_socket_sendall(conn, message) while kms_context.bytes_needed > 0: @@ -233,6 +300,8 @@ async def kms_request(self, kms_context: MongoCryptKmsContext) -> None: conn.close() except MongoCryptError: raise # Propagate MongoCryptError errors directly. + except ConfigurationError: + raise # A callback contract violation is not transient. except Exception as exc: remaining = _csot.remaining() if isinstance(exc, NetworkTimeout) or (remaining is not None and remaining <= 0): @@ -596,6 +665,7 @@ def __init__( codec_options: CodecOptions[_DocumentTypeArg], kms_tls_options: Optional[Mapping[str, Any]] = None, key_expiration_ms: Optional[int] = None, + kms_connect_callback: Optional[AsyncKMSConnectCallback] = None, ) -> None: """Explicit client-side field level encryption. @@ -665,7 +735,18 @@ def __init__( :param key_expiration_ms: The cache expiration time for data encryption keys. Defaults to ``None`` which defers to libmongocrypt's default which is currently 60000. Set to 0 to disable key expiration. - + :param kms_connect_callback: A callable that opens the connection to a + KMS host, used to route KMS requests through an HTTP proxy. It + receives a :class:`~pymongo.encryption_options.KMSConnectContext` + and returns a connected, unwrapped :class:`socket.socket`; the + driver then performs the KMS TLS handshake over it. For an ordinary + HTTP proxy, pass + :class:`~pymongo.encryption_options.AsyncHTTPProxyKMSConnect`. + Defaults to ``None``, meaning the driver connects to KMS hosts + directly. + + .. versionchanged:: 4.18 + Added the `kms_connect_callback` parameter. .. versionchanged:: 4.12 Added the `key_expiration_ms` parameter. .. versionchanged:: 4.0 @@ -709,6 +790,7 @@ def __init__( key_vault_namespace, kms_tls_options=kms_tls_options, key_expiration_ms=key_expiration_ms, + kms_connect_callback=kms_connect_callback, ) self._kms_ssl_contexts = _parse_kms_tls_options(opts._kms_tls_options, _IS_SYNC) self._io_callbacks: Optional[_EncryptionIO] = _EncryptionIO( diff --git a/pymongo/encryption_options.py b/pymongo/encryption_options.py index f2fcd47c65..5e2a575f29 100644 --- a/pymongo/encryption_options.py +++ b/pymongo/encryption_options.py @@ -19,8 +19,15 @@ from __future__ import annotations -from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Optional, TypedDict +import asyncio +import functools +import socket +import ssl +import threading +import time +from collections.abc import Awaitable, Mapping +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Callable, Optional, TypedDict from pymongo.uri_parser_shared import _parse_kms_tls_options @@ -54,6 +61,205 @@ def check_min_pymongocrypt() -> None: ) +@dataclass(frozen=True) +class KMSConnectContext: + """Information about a pending KMS connection. + + Passed to ``kms_connect_callback``, which must return a plain, unwrapped + :class:`socket.socket`. The driver performs the KMS TLS handshake over it + against ``host``, not the peer actually reached, which is what makes + proxying safe. + + 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 left in the timeout budget, else the default KMS + connect timeout. + + .. note:: ``timeoutMS`` does not constrain KMS requests for explicit + encryption, so ``timeout`` is always the default there. Automatic + encryption passes the remaining budget. This deviates from the Client + Side Operations Timeout specification; see PYTHON-6037. + + .. versionadded:: 4.18 + """ + + host: str + port: int + timeout: float + + +# A callback that opens a connection to a KMS host. +AsyncKMSConnectCallback = Callable[[KMSConnectContext], Awaitable[socket.socket]] +KMSConnectCallback = Callable[[KMSConnectContext], socket.socket] + +# Largest CONNECT response header accepted, so a proxy that never sends the +# terminator cannot grow the buffer without bound. +_MAX_CONNECT_HEADER = 8192 + + +def _close_completed_socket(future: asyncio.Future[socket.socket]) -> None: + """Close a socket produced after its awaiting task was cancelled.""" + if not future.cancelled() and future.exception() is None: + future.result().close() + + +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 HTTPProxyKMSConnect + + opts = AutoEncryptionOpts( + kms_providers={"aws": aws_creds}, + key_vault_namespace="keyvault.datakeys", + kms_connect_callback=HTTPProxyKMSConnect("proxy.example.com", 8080), + ) + + To reach the proxy over TLS, pass an :class:`ssl.SSLContext`. 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("proxy.example.com", 8443, proxy_tls) + + Use :class:`AsyncHTTPProxyKMSConnect` with the asynchronous API. + + :param host: Hostname of the proxy. + :param port: Port of the proxy. + :param ssl_context: Optional :class:`ssl.SSLContext` for connecting to the + proxy over TLS. Defaults to ``None``, meaning a plain connection. + + .. versionadded:: 4.18 + """ + + def __init__(self, host: str, port: int, ssl_context: Optional[ssl.SSLContext] = None): + self.host = host + self.port = port + self.ssl_context = ssl_context + + def _tunnel(self, sock: socket.socket, context: KMSConnectContext) -> 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}" + sock.sendall(f"CONNECT {target} HTTP/1.1\r\nHost: {target}\r\n\r\n".encode()) + # A byte at a time: a bulk read could consume tunnelled bytes sent in + # the same segment as the response, and the driver reads those from + # this same socket. + response = bytearray() + while not response.endswith(b"\r\n\r\n"): + 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] + if not status.startswith(b"HTTP/1.1 200"): + 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. Threads rather than tasks, even in + :class:`AsyncHTTPProxyKMSConnect`, because the event loop cannot read + an :class:`ssl.SSLSocket`. + """ + driver_side, relay_side = socket.socketpair() + + def relay(src: socket.socket, dst: socket.socket) -> None: + try: + while True: + buf = src.recv(16384) + if not buf: + break + dst.sendall(buf) + except OSError: + pass + finally: + # EOF the peer instead of closing a socket it may be reading. + try: + dst.shutdown(socket.SHUT_RDWR) + except OSError: + pass + src.close() + + started = [] + try: + for pair in ((relay_side, proxy), (proxy, relay_side)): + thread = threading.Thread(target=relay, args=pair, daemon=True) + thread.start() + started.append(thread) + except BaseException: + # Unblock any thread that did start, then drop every socket. + for sock in (proxy, relay_side, driver_side): + try: + sock.shutdown(socket.SHUT_RDWR) + except OSError: + pass + sock.close() + raise + return driver_side + + def __call__(self, context: KMSConnectContext) -> socket.socket: + # One deadline for all three phases; a timeout per phase would let the + # total run to several times the caller's budget. + deadline = time.monotonic() + context.timeout + sock = socket.create_connection((self.host, self.port), timeout=_remaining(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) + except BaseException: + sock.close() + raise + if self.ssl_context is None: + return sock + try: + return self._bridge(sock) + except BaseException: + sock.close() + raise + + +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.18 + """ + + async def __call__(self, context: KMSConnectContext) -> socket.socket: # type: ignore[override] + # run_in_executor, as auth_oidc.py does for user callbacks. + connect = functools.partial(super().__call__, context) + future = asyncio.get_running_loop().run_in_executor(None, connect) + try: + return await asyncio.shield(future) + except asyncio.CancelledError: + # The thread runs on regardless, so close the socket it returns. + future.add_done_callback(_close_completed_socket) + raise + + class AutoEncryptionOpts: """Options to configure automatic client-side field level encryption.""" @@ -74,6 +280,7 @@ def __init__( bypass_query_analysis: bool = False, encrypted_fields_map: Optional[Mapping[str, Any]] = None, key_expiration_ms: Optional[int] = None, + kms_connect_callback: Optional[Callable[[KMSConnectContext], Any]] = None, ) -> None: """Options to configure automatic client-side field level encryption. @@ -211,7 +418,20 @@ def __init__( :param key_expiration_ms: The cache expiration time for data encryption keys. Defaults to ``None`` which defers to libmongocrypt's default which is currently 60000. Set to 0 to disable key expiration. - + :param kms_connect_callback: A callable that opens the connection to a + KMS host, used to route KMS requests through an HTTP proxy. It + receives a :class:`KMSConnectContext` and returns a connected, + unwrapped :class:`socket.socket`; the driver then performs the KMS + TLS handshake over it. Must be a coroutine function for + :class:`~pymongo.asynchronous.mongo_client.AsyncMongoClient` and a + regular function for + :class:`~pymongo.synchronous.mongo_client.MongoClient`. For an + ordinary HTTP proxy, pass :class:`HTTPProxyKMSConnect` or + :class:`AsyncHTTPProxyKMSConnect`. Defaults to ``None``, meaning + the driver connects to KMS hosts directly. + + .. versionchanged:: 4.18 + Added the `kms_connect_callback` parameter. .. versionchanged:: 4.12 Added the `key_expiration_ms` parameter. .. versionchanged:: 4.2 @@ -258,6 +478,11 @@ def __init__( self._async_kms_ssl_contexts: Optional[dict[str, SSLContext]] = None self._bypass_query_analysis = bypass_query_analysis self._key_expiration_ms = key_expiration_ms + if kms_connect_callback is not None and not callable(kms_connect_callback): + raise TypeError( + f"kms_connect_callback must be callable, not {type(kms_connect_callback)}" + ) + self._kms_connect_callback = kms_connect_callback def _kms_ssl_contexts(self, is_sync: bool) -> dict[str, SSLContext]: if is_sync: diff --git a/pymongo/pool_shared.py b/pymongo/pool_shared.py index 410ffd8189..ec529e810c 100644 --- a/pymongo/pool_shared.py +++ b/pymongo/pool_shared.py @@ -259,16 +259,19 @@ async def _async_create_connection(address: _Address, options: PoolOptions) -> s raise OSError("getaddrinfo failed") -async def _async_configured_socket( - address: _Address, options: PoolOptions +async def _async_wrap_socket_tls( + sock: socket.socket, address: _Address, options: PoolOptions ) -> Union[socket.socket, _sslConn]: - """Given (host, port) and PoolOptions, return a raw configured socket. + """Given a connected socket, (host, port), and PoolOptions, apply TLS. + + The handshake, SNI, and certificate/hostname verification all target + ``address``, which may differ from the peer ``sock`` is connected to, for + example when ``sock`` tunnels through an HTTP proxy. Can raise socket.error, ConnectionFailure, or _CertificateError. - Sets socket's SSL and timeout options. + Sets the socket's SSL and timeout options. """ - sock = await _async_create_connection(address, options) ssl_context = options._ssl_context if ssl_context is None: @@ -315,6 +318,19 @@ async def _async_configured_socket( return ssl_sock +async def _async_configured_socket( + address: _Address, options: PoolOptions +) -> Union[socket.socket, _sslConn]: + """Given (host, port) and PoolOptions, return a raw configured socket. + + Can raise socket.error, ConnectionFailure, or _CertificateError. + + Sets socket's SSL and timeout options. + """ + sock = await _async_create_connection(address, options) + return await _async_wrap_socket_tls(sock, address, options) + + async def _configured_protocol_interface( address: _Address, options: PoolOptions, @@ -465,14 +481,19 @@ def _create_connection(address: _Address, options: PoolOptions) -> socket.socket raise OSError("getaddrinfo failed") -def _configured_socket(address: _Address, options: PoolOptions) -> Union[socket.socket, _sslConn]: - """Given (host, port) and PoolOptions, return a raw configured socket. +def _wrap_socket_tls( + sock: socket.socket, address: _Address, options: PoolOptions +) -> Union[socket.socket, _sslConn]: + """Given a connected socket, (host, port), and PoolOptions, apply TLS. + + The handshake, SNI, and certificate/hostname verification all target + ``address``, which may differ from the peer ``sock`` is connected to, for + example when ``sock`` tunnels through an HTTP proxy. Can raise socket.error, ConnectionFailure, or _CertificateError. - Sets socket's SSL and timeout options. + Sets the socket's SSL and timeout options. """ - sock = _create_connection(address, options) ssl_context = options._ssl_context if ssl_context is None: @@ -514,6 +535,17 @@ def _configured_socket(address: _Address, options: PoolOptions) -> Union[socket. return ssl_sock +def _configured_socket(address: _Address, options: PoolOptions) -> Union[socket.socket, _sslConn]: + """Given (host, port) and PoolOptions, return a raw configured socket. + + Can raise socket.error, ConnectionFailure, or _CertificateError. + + Sets socket's SSL and timeout options. + """ + sock = _create_connection(address, options) + return _wrap_socket_tls(sock, address, options) + + def _configured_socket_interface( address: _Address, options: PoolOptions, diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index 014d162e2b..723c3dfc22 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -18,7 +18,9 @@ import contextlib import enum +import inspect import socket +import ssl import time as time # noqa: PLC0414 # needed in sync version import uuid import weakref @@ -59,6 +61,8 @@ from pymongo.daemon import _spawn_daemon from pymongo.encryption_options import ( AutoEncryptionOpts, + KMSConnectCallback, + KMSConnectContext, RangeOpts, TextOpts, check_min_pymongocrypt, @@ -78,6 +82,7 @@ from pymongo.pool_shared import ( _configured_socket, _raise_connection_failure, + _wrap_socket_tls, ) from pymongo.read_concern import ReadConcern from pymongo.results import BulkWriteResult, DeleteResult @@ -111,9 +116,65 @@ _KEY_VAULT_OPTS = CodecOptions(document_class=RawBSONDocument) -def _connect_kms(address: _Address, opts: PoolOptions) -> Union[socket.socket, _sslConn]: +def _close_rejected_kms_socket(obj: Any) -> None: + """Close a rejected kms_connect_callback return value, if it can be closed. + + The caller may have handed us a live socket, and nothing else will close it: + _connect_kms raises before its result reaches the caller's ``finally``. The + value can be anything a user returned, so closing is strictly best effort. + """ + close = getattr(obj, "close", None) + if callable(close): + with contextlib.suppress(Exception): + close() + + +def _connect_kms( + address: _Address, + opts: PoolOptions, + kms_connect_callback: Optional[KMSConnectCallback], + timeout: float, +) -> Union[socket.socket, _sslConn]: + if kms_connect_callback is None: + try: + return _configured_socket(address, opts) + except Exception as exc: + _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) + + # TLS targets address, not the peer, so verification follows the KMS host. + result = kms_connect_callback( + KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) + ) + # The synchronous module takes a regular function and awaits nothing. + if not _IS_SYNC and not inspect.isawaitable(result): + _close_rejected_kms_socket(result) + raise ConfigurationError( + "kms_connect_callback must be a coroutine function for the async " + f"API, but returned {type(result)}." + ) + sock = result + if not isinstance(sock, socket.socket) or isinstance(sock, ssl.SSLSocket): + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return a connected, unwrapped " + f"socket.socket, not {type(sock)}; consider HTTPProxyKMSConnect." + ) + # wrap_socket refuses a non-blocking socket, so normalize the mode here. try: - return _configured_socket(address, opts) + sock.getpeername() + except OSError: + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return an already connected socket." + ) from None + if sock.getsockopt(socket.SOL_SOCKET, socket.SO_TYPE) != socket.SOCK_STREAM: + _close_rejected_kms_socket(sock) + raise ConfigurationError( + "kms_connect_callback must return a stream socket, not a datagram one." + ) + sock.settimeout(opts.socket_timeout) + try: + return _wrap_socket_tls(sock, address, opts) except Exception as exc: _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) @@ -183,20 +244,26 @@ def kms_request(self, kms_context: MongoCryptKmsContext) -> None: False, # disable_ocsp_endpoint_check _IS_SYNC, ) - # CSOT: set timeout for socket creation. + address = parse_host(endpoint, _HTTPS_PORT) + sleep_u = kms_context.usleep + if sleep_u: + sleep_sec = float(sleep_u) / 1e6 + time.sleep(sleep_sec) + # CSOT: set timeout for socket creation. After the retry backoff above, + # so the budget reflects what the sleep consumed. connect_timeout = max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001) opts = PoolOptions( connect_timeout=connect_timeout, socket_timeout=connect_timeout, ssl_context=ctx, ) - address = parse_host(endpoint, _HTTPS_PORT) - sleep_u = kms_context.usleep - if sleep_u: - sleep_sec = float(sleep_u) / 1e6 - time.sleep(sleep_sec) try: - conn = _connect_kms(address, opts) + conn = _connect_kms( + address, + opts, + self.opts._kms_connect_callback, + connect_timeout, + ) try: sendall(conn, message) while kms_context.bytes_needed > 0: @@ -232,6 +299,8 @@ def kms_request(self, kms_context: MongoCryptKmsContext) -> None: conn.close() except MongoCryptError: raise # Propagate MongoCryptError errors directly. + except ConfigurationError: + raise # A callback contract violation is not transient. except Exception as exc: remaining = _csot.remaining() if isinstance(exc, NetworkTimeout) or (remaining is not None and remaining <= 0): @@ -593,6 +662,7 @@ def __init__( codec_options: CodecOptions[_DocumentTypeArg], kms_tls_options: Optional[Mapping[str, Any]] = None, key_expiration_ms: Optional[int] = None, + kms_connect_callback: Optional[KMSConnectCallback] = None, ) -> None: """Explicit client-side field level encryption. @@ -662,7 +732,18 @@ def __init__( :param key_expiration_ms: The cache expiration time for data encryption keys. Defaults to ``None`` which defers to libmongocrypt's default which is currently 60000. Set to 0 to disable key expiration. - + :param kms_connect_callback: A callable that opens the connection to a + KMS host, used to route KMS requests through an HTTP proxy. It + receives a :class:`~pymongo.encryption_options.KMSConnectContext` + and returns a connected, unwrapped :class:`socket.socket`; the + driver then performs the KMS TLS handshake over it. For an ordinary + HTTP proxy, pass + :class:`~pymongo.encryption_options.HTTPProxyKMSConnect`. + Defaults to ``None``, meaning the driver connects to KMS hosts + directly. + + .. versionchanged:: 4.18 + Added the `kms_connect_callback` parameter. .. versionchanged:: 4.12 Added the `key_expiration_ms` parameter. .. versionchanged:: 4.0 @@ -702,6 +783,7 @@ def __init__( key_vault_namespace, kms_tls_options=kms_tls_options, key_expiration_ms=key_expiration_ms, + kms_connect_callback=kms_connect_callback, ) self._kms_ssl_contexts = _parse_kms_tls_options(opts._kms_tls_options, _IS_SYNC) self._io_callbacks: Optional[_EncryptionIO] = _EncryptionIO( diff --git a/test/asynchronous/test_encryption.py b/test/asynchronous/test_encryption.py index daeb18607a..6860749b64 100644 --- a/test/asynchronous/test_encryption.py +++ b/test/asynchronous/test_encryption.py @@ -16,8 +16,10 @@ from __future__ import annotations +import asyncio import base64 import copy +import dataclasses import http.client import json import os @@ -28,12 +30,16 @@ import ssl import sys import textwrap +import threading +import time import traceback import uuid import warnings +from asyncio.trsock import TransportSocket from collections.abc import Mapping from threading import Thread from typing import Any, Optional +from unittest import mock import pytest @@ -59,11 +65,26 @@ from bson.son import SON from pymongo import ReadPreference from pymongo.asynchronous import encryption -from pymongo.asynchronous.encryption import Algorithm, AsyncClientEncryption, QueryType +from pymongo.asynchronous.encryption import ( + Algorithm, + AsyncClientEncryption, + QueryType, + _connect_kms, + _EncryptionIO, + _wrap_encryption_errors, +) from pymongo.asynchronous.helpers import anext from pymongo.asynchronous.mongo_client import AsyncMongoClient from pymongo.cursor_shared import CursorType -from pymongo.encryption_options import _HAVE_PYMONGOCRYPT, AutoEncryptionOpts, RangeOpts, TextOpts +from pymongo.encryption_options import ( + _HAVE_PYMONGOCRYPT, + AsyncHTTPProxyKMSConnect, + AutoEncryptionOpts, + HTTPProxyKMSConnect, + KMSConnectContext, + RangeOpts, + TextOpts, +) from pymongo.errors import ( AutoReconnect, BulkWriteError, @@ -78,6 +99,8 @@ WriteError, ) from pymongo.operations import InsertOne, ReplaceOne, UpdateOne +from pymongo.pool_options import PoolOptions +from pymongo.ssl_support import get_ssl_context from pymongo.write_concern import WriteConcern from test import ( unittest, @@ -213,6 +236,471 @@ async def test_init_kms_tls_options(self): self.assertEqual(ctx.check_hostname, True) self.assertEqual(ctx.verify_mode, ssl.CERT_REQUIRED) + @unittest.skipUnless(_HAVE_PYMONGOCRYPT, "pymongocrypt is not installed") + async def test_init_kms_connect_callback(self): + opts = AutoEncryptionOpts({}, "k.d") + self.assertIsNone(opts._kms_connect_callback) + + async def callback(context): + raise AssertionError("not called") + + opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) + self.assertIs(opts._kms_connect_callback, callback) + + for bad in [1, "not-callable", object()]: + with self.assertRaisesRegex(TypeError, "kms_connect_callback must be callable"): + AutoEncryptionOpts({}, "k.d", kms_connect_callback=bad) # type: ignore[arg-type] + + context = KMSConnectContext(host="kms.example.com", port=443, timeout=9.5) + self.assertEqual(context.host, "kms.example.com") + self.assertEqual(context.port, 443) + self.assertEqual(context.timeout, 9.5) + with self.assertRaises(dataclasses.FrozenInstanceError): + context.host = "evil.example.com" # type: ignore[misc] + + +class TestKmsConnectCallbackUnit(AsyncPyMongoTestCase): + """Contract checks for kms_connect_callback that need no KMS server.""" + + @staticmethod + def _pool_options(): + return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=None) + + async def test_non_socket_return_raises_configuration_error(self): + async def callback(context): + return "not-a-socket" + + with self.assertRaisesRegex(ConfigurationError, "must return a connected"): + await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + async def test_already_wrapped_socket_is_rejected(self): + # ssl.SSLSocket passes isinstance but cannot be TLS-wrapped again. + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + left, right = socket.socketpair() + self.addCleanup(right.close) + # No peer needed to produce a genuine ssl.SSLSocket. + wrapped = ctx.wrap_socket(left, do_handshake_on_connect=False, server_hostname="x") + self.addCleanup(wrapped.close) + + async def callback(context): + return wrapped + + with self.assertRaisesRegex(ConfigurationError, "unwrapped"): + await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + async def test_context_receives_host_port_and_timeout(self): + received = [] + left, right = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + + async def callback(context): + received.append(context) + return left + + # ssl_context=None returns the socket unchanged, so a plain socket is accepted. + conn = await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 12.5) + self.assertIs(conn, left) + + self.assertEqual(len(received), 1) + self.assertEqual(received[0].host, "kms.example.com") + self.assertEqual(received[0].port, 443) + self.assertEqual(received[0].timeout, 12.5) + + async def test_non_blocking_socket_from_callback_is_accepted(self): + # Without the driver normalizing the mode, this raises ValueError. + server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_ctx.load_cert_chain(CLIENT_PEM) + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def serve(): + try: + conn, _ = listener.accept() + server_ctx.wrap_socket(conn, server_side=True).close() + except OSError: + pass + + threading.Thread(target=serve, daemon=True).start() + + # Built as the driver does, for the flavor-correct type; the local cert won't verify. + client_ctx = get_ssl_context(None, None, None, None, True, True, False, _IS_SYNC) + options = PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=client_ctx) + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + sock.setblocking(False) + return sock + + async def callback(context): + if _IS_SYNC: + return connect() + return await asyncio.get_running_loop().run_in_executor(None, connect) + + conn = await _connect_kms(listener.getsockname(), options, callback, 10.0) + self.addCleanup(conn.close) + self.assertIsNotNone(conn.gettimeout()) + + async def test_asyncio_transport_socket_is_rejected(self): + # get_extra_info("socket") is a TransportSocket, not a socket.socket. + left, right = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + + async def callback(context): + return TransportSocket(left) + + with self.assertRaisesRegex(ConfigurationError, "TransportSocket"): + await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + async def test_http_proxy_helper_tunnels_and_reports_refusal(self): + # Covers the CONNECT handshake without KMS credentials. + accepted = [] + + def stub(listener, reply): + try: + conn, _ = listener.accept() + except OSError: + return + accepted.append(conn.recv(4096)) + conn.sendall(reply) + conn.close() + + def run_stub(reply): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + threading.Thread(target=stub, args=(listener, reply), daemon=True).start() + return listener.getsockname() + + host, port = run_stub(b"HTTP/1.1 200 Connection Established\r\n\r\n") + callback = AsyncHTTPProxyKMSConnect(host, port) + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + sock = await callback(context) + self.addCleanup(sock.close) + self.assertIsInstance(sock, socket.socket) + self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") + + host, port = run_stub(b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n") + with self.assertRaisesRegex(OSError, "refused CONNECT"): + await AsyncHTTPProxyKMSConnect(host, port)(context) + + async def test_tls_proxy_helper_bridges_the_tunnel(self): + # Covers the TLS-proxy path and the socketpair relay without KMS creds. + server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_ctx.load_cert_chain(CLIENT_PEM) + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + conn = None + try: + conn, _ = listener.accept() + tls = server_ctx.wrap_socket(conn, server_side=True) + request = b"" + while b"\r\n\r\n" not in request: + chunk = tls.recv(4096) + if not chunk: + return + request += chunk + tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + # The tunnelled peer speaks only after the client does, as a + # TLS server would. + tls.sendall(b"echo:" + tls.recv(64)) + tls.close() + except OSError: + pass + finally: + if conn is not None: + conn.close() + + threading.Thread(target=stub_proxy, daemon=True).start() + + client_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + client_ctx.check_hostname = False + client_ctx.verify_mode = ssl.CERT_NONE + host, port = listener.getsockname() + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + + sock = await AsyncHTTPProxyKMSConnect(host, port, client_ctx)(context) + self.addCleanup(sock.close) + sock.settimeout(10) + sock.sendall(b"ping") + self.assertEqual(sock.recv(64), b"echo:ping") + + async def test_non_coroutine_callback_is_rejected(self): + # The async API needs a coroutine function; a plain def must not be + # awaited and retried. + if _IS_SYNC: + raise unittest.SkipTest("a regular function is correct for the sync API") + + left, right = socket.socketpair() + self.addCleanup(right.close) + + def callback(context): + return left + + with self.assertRaisesRegex(ConfigurationError, "coroutine function"): + await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + async def test_proxy_closing_before_connect_reply_raises(self): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + try: + conn, _ = listener.accept() + # Read the CONNECT request, then hang up without replying. + conn.recv(4096) + conn.close() + except OSError: + pass + + threading.Thread(target=stub_proxy, daemon=True).start() + + host, port = listener.getsockname() + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + with self.assertRaisesRegex(OSError, "proxy closed the connection"): + await AsyncHTTPProxyKMSConnect(host, port)(context) + + async def test_tunnel_keeps_bytes_sent_with_the_connect_reply(self): + # A proxy may coalesce its 200 with tunnelled bytes; reading past the + # header would silently drop them. + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + try: + conn, _ = listener.accept() + conn.recv(4096) + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\nearly-bytes") + conn.close() + except OSError: + pass + + threading.Thread(target=stub_proxy, daemon=True).start() + + host, port = listener.getsockname() + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + sock = await AsyncHTTPProxyKMSConnect(host, port)(context) + self.addCleanup(sock.close) + sock.settimeout(10) + self.assertEqual(sock.recv(64), b"early-bytes") + + async def test_unconnected_socket_from_callback_is_rejected(self): + # The contract says connected; an unconnected socket would otherwise + # fail later as a transient error and be retried. + bare = socket.socket() + self.addCleanup(bare.close) + + async def callback(context): + return bare + + with self.assertRaisesRegex(ConfigurationError, "already connected"): + await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + async def test_ipv6_host_is_bracketed_in_connect(self): + accepted = [] + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + try: + conn, _ = listener.accept() + accepted.append(conn.recv(4096)) + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + conn.close() + except OSError: + pass + + threading.Thread(target=stub_proxy, daemon=True).start() + + host, port = listener.getsockname() + context = KMSConnectContext(host="::1", port=443, timeout=10) + sock = await AsyncHTTPProxyKMSConnect(host, port)(context) + self.addCleanup(sock.close) + self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT [::1]:443 HTTP/1.1") + + async def test_oversized_connect_response_is_rejected(self): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + conn = None + try: + conn, _ = listener.accept() + conn.recv(4096) + # Never sends the terminator. + while True: + conn.sendall(b"x" * 1024) + except OSError: + pass + finally: + if conn is not None: + conn.close() + + threading.Thread(target=stub_proxy, daemon=True).start() + + host, port = listener.getsockname() + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + with self.assertRaisesRegex(OSError, "oversized CONNECT response"): + await AsyncHTTPProxyKMSConnect(host, port)(context) + + async def test_remaining_raises_once_the_deadline_passes(self): + from pymongo.encryption_options import _remaining + + self.assertGreater(_remaining(time.monotonic() + 5), 0) + with self.assertRaises(socket.timeout): + _remaining(time.monotonic() - 1) + + async def test_datagram_socket_from_callback_is_rejected(self): + # A connected UDP socket passes isinstance and getpeername, but TLS + # then raises NotImplementedError, which would be retried. + left = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + right = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + self.addCleanup(left.close) + self.addCleanup(right.close) + right.bind(("127.0.0.1", 0)) + left.connect(right.getsockname()) + + async def callback(context): + return left + + with self.assertRaisesRegex(ConfigurationError, "stream socket"): + await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + async def test_kms_request_does_not_retry_a_contract_violation(self): + # _connect_kms has no retry loop; the no-retry guarantee is in + # kms_request, so exercise that instead. + calls = [] + + async def callback(context): + calls.append(context) + return "not-a-socket" + + opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) + io = _EncryptionIO(None, mock.MagicMock(), None, opts) + + class StubKmsContext: + endpoint = "kms.example.com:443" + message = b"request" + kms_provider = "aws" + usleep = 0 + bytes_needed = 1 + + def feed(self, data): + raise AssertionError("should not reach the socket") + + def fail(self): + raise AssertionError("a contract violation must not be retried") + + with self.assertRaises(ConfigurationError): + await io.kms_request(StubKmsContext()) + self.assertEqual(len(calls), 1) + + async def test_contract_violation_surfaces_as_encryption_error(self): + # Public operations run under _wrap_encryption_errors, so callers see + # EncryptionError with ConfigurationError as its cause. + with self.assertRaises(EncryptionError) as caught: + with _wrap_encryption_errors(): + raise ConfigurationError("kms_connect_callback must return ...") + self.assertIsInstance(caught.exception.__cause__, ConfigurationError) + + async def test_bridge_failure_closes_the_proxy_socket(self): + # A failure inside _bridge must not strand the connected proxy socket. + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_ctx.load_cert_chain(CLIENT_PEM) + + def stub_proxy(): + conn = None + try: + conn, _ = listener.accept() + 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() + except OSError: + pass + finally: + if conn is not None: + conn.close() + + threading.Thread(target=stub_proxy, daemon=True).start() + + captured = [] + + def failing_bridge(self, proxy): + captured.append(proxy) + raise OSError("no file descriptors") + + host, port = listener.getsockname() + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + + with mock.patch.object(HTTPProxyKMSConnect, "_bridge", failing_bridge): + with self.assertRaisesRegex(OSError, "no file descriptors"): + await AsyncHTTPProxyKMSConnect(host, port, ctx)(context) + + self.assertEqual(captured[0].fileno(), -1, "proxy socket was left open") + + async def test_network_error_from_callback_propagates(self): + async def callback(context): + raise OSError("proxy unreachable") + + # Not a ConfigurationError, so kms_request retries it. + with self.assertRaises(OSError): + await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + @unittest.skipUnless(_HAVE_PYMONGOCRYPT, "pymongocrypt is not installed") + async def test_client_encryption_accepts_callback(self): + async def callback(context): + raise AssertionError("not called") + + client = self.simple_client() + encryption = AsyncClientEncryption( + {"local": {"key": b"\x00" * 96}}, + "keyvault.datakeys", + client, + OPTS, + kms_connect_callback=callback, + ) + self.addAsyncCleanup(encryption.close) + self.assertIs(encryption._io_callbacks.opts._kms_connect_callback, callback) + + @unittest.skipUnless(_HAVE_PYMONGOCRYPT, "pymongocrypt is not installed") + async def test_client_encryption_rejects_non_callable(self): + client = self.simple_client() + with self.assertRaisesRegex(TypeError, "kms_connect_callback must be callable"): + AsyncClientEncryption( + {"local": {"key": b"\x00" * 96}}, + "keyvault.datakeys", + client, + OPTS, + kms_connect_callback="not-callable", # type: ignore[arg-type] + ) + class TestClientOptions(AsyncPyMongoTestCase): async def test_default(self): @@ -252,9 +740,15 @@ def create_client_encryption( key_vault_client: AsyncMongoClient, codec_options: CodecOptions, kms_tls_options: Optional[Mapping[str, Any]] = None, + kms_connect_callback: Optional[Any] = None, ): client_encryption = AsyncClientEncryption( - kms_providers, key_vault_namespace, key_vault_client, codec_options, kms_tls_options + kms_providers, + key_vault_namespace, + key_vault_client, + codec_options, + kms_tls_options, + kms_connect_callback=kms_connect_callback, ) self.addAsyncCleanup(client_encryption.close) return client_encryption @@ -267,9 +761,15 @@ def unmanaged_create_client_encryption( key_vault_client: AsyncMongoClient, codec_options: CodecOptions, kms_tls_options: Optional[Mapping[str, Any]] = None, + kms_connect_callback: Optional[Any] = None, ): client_encryption = AsyncClientEncryption( - kms_providers, key_vault_namespace, key_vault_client, codec_options, kms_tls_options + kms_providers, + key_vault_namespace, + key_vault_client, + codec_options, + kms_tls_options, + kms_connect_callback=kms_connect_callback, ) return client_encryption @@ -1916,6 +2416,188 @@ async def test_invalid_hostname_in_kms_certificate(self): await self.client_encrypted.create_data_key("aws", master_key=key) +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", +} + + +# https://github.com/mongodb/specifications/blob/master/source/client-side-encryption/tests/README.md#kms-connect-callback +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(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) + + async def tls_callback(self, context): + self.callback_calls.append(context) + ctx = ssl.create_default_context(cafile=CA_PEM) + ctx.check_hostname = False + callback = AsyncHTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_TLS_PROXY_PORT, ctx) + 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: + ctx = ssl.create_default_context(cafile=CA_PEM) + ctx.check_hostname = False + conn = http.client.HTTPSConnection( + f"{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", context=ctx + ) + 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) + + self.assertGreaterEqual(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(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) + + # https://github.com/mongodb/specifications/blob/master/source/client-side-encryption/tests/README.md#kms-tls-options-tests class TestKmsTLSOptions(AsyncEncryptionIntegrationTest): @unittest.skipUnless(any(AWS_CREDS.values()), "AWS environment credentials are not set") diff --git a/test/asynchronous/test_pooling.py b/test/asynchronous/test_pooling.py index 063f5f06ec..5b97fb108f 100644 --- a/test/asynchronous/test_pooling.py +++ b/test/asynchronous/test_pooling.py @@ -34,13 +34,19 @@ from pymongo.hello import HelloCompat from pymongo.lock import _async_create_lock from pymongo.monitoring import _EventListeners +from pymongo.pool_shared import _async_wrap_socket_tls from test.asynchronous.utils import async_get_pool, async_joinall, flaky sys.path[0:0] = [""] from pymongo.asynchronous.pool import Pool, PoolOptions from pymongo.socket_checker import SocketChecker -from test.asynchronous import AsyncIntegrationTest, async_client_context, unittest +from test.asynchronous import ( + AsyncIntegrationTest, + AsyncPyMongoTestCase, + async_client_context, + unittest, +) from test.asynchronous.helpers import ConcurrentRunner from test.utils_shared import CMAPListener, delay @@ -759,5 +765,18 @@ def test_certificate_error_is_not_labeled_overloaded(self): self.assertFalse(err.has_error_label("SystemOverloadedError")) +class TestWrapSocketTLS(AsyncPyMongoTestCase): + async def test_wrap_socket_tls_without_ssl_context_returns_same_socket(self): + options = PoolOptions(socket_timeout=7.5) + left, right = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + + result = await _async_wrap_socket_tls(left, ("kms.example.com", 443), options) + + self.assertIs(result, left) + self.assertEqual(result.gettimeout(), 7.5) + + if __name__ == "__main__": unittest.main() diff --git a/test/test_encryption.py b/test/test_encryption.py index 744db01b1b..2325f1e0f7 100644 --- a/test/test_encryption.py +++ b/test/test_encryption.py @@ -16,8 +16,10 @@ from __future__ import annotations +import asyncio import base64 import copy +import dataclasses import http.client import json import os @@ -28,12 +30,16 @@ import ssl import sys import textwrap +import threading +import time import traceback import uuid import warnings +from asyncio.trsock import TransportSocket from collections.abc import Mapping from threading import Thread from typing import Any, Optional +from unittest import mock import pytest @@ -59,7 +65,14 @@ from bson.son import SON from pymongo import ReadPreference from pymongo.cursor_shared import CursorType -from pymongo.encryption_options import _HAVE_PYMONGOCRYPT, AutoEncryptionOpts, RangeOpts, TextOpts +from pymongo.encryption_options import ( + _HAVE_PYMONGOCRYPT, + AutoEncryptionOpts, + HTTPProxyKMSConnect, + KMSConnectContext, + RangeOpts, + TextOpts, +) from pymongo.errors import ( AutoReconnect, BulkWriteError, @@ -74,8 +87,17 @@ WriteError, ) from pymongo.operations import InsertOne, ReplaceOne, UpdateOne +from pymongo.pool_options import PoolOptions +from pymongo.ssl_support import get_ssl_context from pymongo.synchronous import encryption -from pymongo.synchronous.encryption import Algorithm, ClientEncryption, QueryType +from pymongo.synchronous.encryption import ( + Algorithm, + ClientEncryption, + QueryType, + _connect_kms, + _EncryptionIO, + _wrap_encryption_errors, +) from pymongo.synchronous.helpers import next from pymongo.synchronous.mongo_client import MongoClient from pymongo.write_concern import WriteConcern @@ -213,6 +235,471 @@ def test_init_kms_tls_options(self): self.assertEqual(ctx.check_hostname, True) self.assertEqual(ctx.verify_mode, ssl.CERT_REQUIRED) + @unittest.skipUnless(_HAVE_PYMONGOCRYPT, "pymongocrypt is not installed") + def test_init_kms_connect_callback(self): + opts = AutoEncryptionOpts({}, "k.d") + self.assertIsNone(opts._kms_connect_callback) + + def callback(context): + raise AssertionError("not called") + + opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) + self.assertIs(opts._kms_connect_callback, callback) + + for bad in [1, "not-callable", object()]: + with self.assertRaisesRegex(TypeError, "kms_connect_callback must be callable"): + AutoEncryptionOpts({}, "k.d", kms_connect_callback=bad) # type: ignore[arg-type] + + context = KMSConnectContext(host="kms.example.com", port=443, timeout=9.5) + self.assertEqual(context.host, "kms.example.com") + self.assertEqual(context.port, 443) + self.assertEqual(context.timeout, 9.5) + with self.assertRaises(dataclasses.FrozenInstanceError): + context.host = "evil.example.com" # type: ignore[misc] + + +class TestKmsConnectCallbackUnit(PyMongoTestCase): + """Contract checks for kms_connect_callback that need no KMS server.""" + + @staticmethod + def _pool_options(): + return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=None) + + def test_non_socket_return_raises_configuration_error(self): + def callback(context): + return "not-a-socket" + + with self.assertRaisesRegex(ConfigurationError, "must return a connected"): + _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + def test_already_wrapped_socket_is_rejected(self): + # ssl.SSLSocket passes isinstance but cannot be TLS-wrapped again. + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + left, right = socket.socketpair() + self.addCleanup(right.close) + # No peer needed to produce a genuine ssl.SSLSocket. + wrapped = ctx.wrap_socket(left, do_handshake_on_connect=False, server_hostname="x") + self.addCleanup(wrapped.close) + + def callback(context): + return wrapped + + with self.assertRaisesRegex(ConfigurationError, "unwrapped"): + _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + def test_context_receives_host_port_and_timeout(self): + received = [] + left, right = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + + def callback(context): + received.append(context) + return left + + # ssl_context=None returns the socket unchanged, so a plain socket is accepted. + conn = _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 12.5) + self.assertIs(conn, left) + + self.assertEqual(len(received), 1) + self.assertEqual(received[0].host, "kms.example.com") + self.assertEqual(received[0].port, 443) + self.assertEqual(received[0].timeout, 12.5) + + def test_non_blocking_socket_from_callback_is_accepted(self): + # Without the driver normalizing the mode, this raises ValueError. + server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_ctx.load_cert_chain(CLIENT_PEM) + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def serve(): + try: + conn, _ = listener.accept() + server_ctx.wrap_socket(conn, server_side=True).close() + except OSError: + pass + + threading.Thread(target=serve, daemon=True).start() + + # Built as the driver does, for the flavor-correct type; the local cert won't verify. + client_ctx = get_ssl_context(None, None, None, None, True, True, False, _IS_SYNC) + options = PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=client_ctx) + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + sock.setblocking(False) + return sock + + def callback(context): + if _IS_SYNC: + return connect() + return asyncio.get_running_loop().run_in_executor(None, connect) + + conn = _connect_kms(listener.getsockname(), options, callback, 10.0) + self.addCleanup(conn.close) + self.assertIsNotNone(conn.gettimeout()) + + def test_asyncio_transport_socket_is_rejected(self): + # get_extra_info("socket") is a TransportSocket, not a socket.socket. + left, right = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + + def callback(context): + return TransportSocket(left) + + with self.assertRaisesRegex(ConfigurationError, "TransportSocket"): + _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + def test_http_proxy_helper_tunnels_and_reports_refusal(self): + # Covers the CONNECT handshake without KMS credentials. + accepted = [] + + def stub(listener, reply): + try: + conn, _ = listener.accept() + except OSError: + return + accepted.append(conn.recv(4096)) + conn.sendall(reply) + conn.close() + + def run_stub(reply): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + threading.Thread(target=stub, args=(listener, reply), daemon=True).start() + return listener.getsockname() + + host, port = run_stub(b"HTTP/1.1 200 Connection Established\r\n\r\n") + callback = HTTPProxyKMSConnect(host, port) + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + sock = callback(context) + self.addCleanup(sock.close) + self.assertIsInstance(sock, socket.socket) + self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT kms.example.com:443 HTTP/1.1") + + host, port = run_stub(b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n") + with self.assertRaisesRegex(OSError, "refused CONNECT"): + HTTPProxyKMSConnect(host, port)(context) + + def test_tls_proxy_helper_bridges_the_tunnel(self): + # Covers the TLS-proxy path and the socketpair relay without KMS creds. + server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_ctx.load_cert_chain(CLIENT_PEM) + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + conn = None + try: + conn, _ = listener.accept() + tls = server_ctx.wrap_socket(conn, server_side=True) + request = b"" + while b"\r\n\r\n" not in request: + chunk = tls.recv(4096) + if not chunk: + return + request += chunk + tls.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + # The tunnelled peer speaks only after the client does, as a + # TLS server would. + tls.sendall(b"echo:" + tls.recv(64)) + tls.close() + except OSError: + pass + finally: + if conn is not None: + conn.close() + + threading.Thread(target=stub_proxy, daemon=True).start() + + client_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + client_ctx.check_hostname = False + client_ctx.verify_mode = ssl.CERT_NONE + host, port = listener.getsockname() + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + + sock = HTTPProxyKMSConnect(host, port, client_ctx)(context) + self.addCleanup(sock.close) + sock.settimeout(10) + sock.sendall(b"ping") + self.assertEqual(sock.recv(64), b"echo:ping") + + def test_non_coroutine_callback_is_rejected(self): + # The async API needs a coroutine function; a plain def must not be + # awaited and retried. + if _IS_SYNC: + raise unittest.SkipTest("a regular function is correct for the sync API") + + left, right = socket.socketpair() + self.addCleanup(right.close) + + def callback(context): + return left + + with self.assertRaisesRegex(ConfigurationError, "coroutine function"): + _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + def test_proxy_closing_before_connect_reply_raises(self): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + try: + conn, _ = listener.accept() + # Read the CONNECT request, then hang up without replying. + conn.recv(4096) + conn.close() + except OSError: + pass + + threading.Thread(target=stub_proxy, daemon=True).start() + + host, port = listener.getsockname() + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + with self.assertRaisesRegex(OSError, "proxy closed the connection"): + HTTPProxyKMSConnect(host, port)(context) + + def test_tunnel_keeps_bytes_sent_with_the_connect_reply(self): + # A proxy may coalesce its 200 with tunnelled bytes; reading past the + # header would silently drop them. + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + try: + conn, _ = listener.accept() + conn.recv(4096) + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\nearly-bytes") + conn.close() + except OSError: + pass + + threading.Thread(target=stub_proxy, daemon=True).start() + + host, port = listener.getsockname() + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + sock = HTTPProxyKMSConnect(host, port)(context) + self.addCleanup(sock.close) + sock.settimeout(10) + self.assertEqual(sock.recv(64), b"early-bytes") + + def test_unconnected_socket_from_callback_is_rejected(self): + # The contract says connected; an unconnected socket would otherwise + # fail later as a transient error and be retried. + bare = socket.socket() + self.addCleanup(bare.close) + + def callback(context): + return bare + + with self.assertRaisesRegex(ConfigurationError, "already connected"): + _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + def test_ipv6_host_is_bracketed_in_connect(self): + accepted = [] + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + try: + conn, _ = listener.accept() + accepted.append(conn.recv(4096)) + conn.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + conn.close() + except OSError: + pass + + threading.Thread(target=stub_proxy, daemon=True).start() + + host, port = listener.getsockname() + context = KMSConnectContext(host="::1", port=443, timeout=10) + sock = HTTPProxyKMSConnect(host, port)(context) + self.addCleanup(sock.close) + self.assertEqual(accepted[0].split(b"\r\n")[0], b"CONNECT [::1]:443 HTTP/1.1") + + def test_oversized_connect_response_is_rejected(self): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + def stub_proxy(): + conn = None + try: + conn, _ = listener.accept() + conn.recv(4096) + # Never sends the terminator. + while True: + conn.sendall(b"x" * 1024) + except OSError: + pass + finally: + if conn is not None: + conn.close() + + threading.Thread(target=stub_proxy, daemon=True).start() + + host, port = listener.getsockname() + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + with self.assertRaisesRegex(OSError, "oversized CONNECT response"): + HTTPProxyKMSConnect(host, port)(context) + + def test_remaining_raises_once_the_deadline_passes(self): + from pymongo.encryption_options import _remaining + + self.assertGreater(_remaining(time.monotonic() + 5), 0) + with self.assertRaises(socket.timeout): + _remaining(time.monotonic() - 1) + + def test_datagram_socket_from_callback_is_rejected(self): + # A connected UDP socket passes isinstance and getpeername, but TLS + # then raises NotImplementedError, which would be retried. + left = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + right = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + self.addCleanup(left.close) + self.addCleanup(right.close) + right.bind(("127.0.0.1", 0)) + left.connect(right.getsockname()) + + def callback(context): + return left + + with self.assertRaisesRegex(ConfigurationError, "stream socket"): + _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + def test_kms_request_does_not_retry_a_contract_violation(self): + # _connect_kms has no retry loop; the no-retry guarantee is in + # kms_request, so exercise that instead. + calls = [] + + def callback(context): + calls.append(context) + return "not-a-socket" + + opts = AutoEncryptionOpts({}, "k.d", kms_connect_callback=callback) + io = _EncryptionIO(None, mock.MagicMock(), None, opts) + + class StubKmsContext: + endpoint = "kms.example.com:443" + message = b"request" + kms_provider = "aws" + usleep = 0 + bytes_needed = 1 + + def feed(self, data): + raise AssertionError("should not reach the socket") + + def fail(self): + raise AssertionError("a contract violation must not be retried") + + with self.assertRaises(ConfigurationError): + io.kms_request(StubKmsContext()) + self.assertEqual(len(calls), 1) + + def test_contract_violation_surfaces_as_encryption_error(self): + # Public operations run under _wrap_encryption_errors, so callers see + # EncryptionError with ConfigurationError as its cause. + with self.assertRaises(EncryptionError) as caught: + with _wrap_encryption_errors(): + raise ConfigurationError("kms_connect_callback must return ...") + self.assertIsInstance(caught.exception.__cause__, ConfigurationError) + + def test_bridge_failure_closes_the_proxy_socket(self): + # A failure inside _bridge must not strand the connected proxy socket. + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + self.addCleanup(listener.close) + + server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_ctx.load_cert_chain(CLIENT_PEM) + + def stub_proxy(): + conn = None + try: + conn, _ = listener.accept() + 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() + except OSError: + pass + finally: + if conn is not None: + conn.close() + + threading.Thread(target=stub_proxy, daemon=True).start() + + captured = [] + + def failing_bridge(self, proxy): + captured.append(proxy) + raise OSError("no file descriptors") + + host, port = listener.getsockname() + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + context = KMSConnectContext(host="kms.example.com", port=443, timeout=10) + + with mock.patch.object(HTTPProxyKMSConnect, "_bridge", failing_bridge): + with self.assertRaisesRegex(OSError, "no file descriptors"): + HTTPProxyKMSConnect(host, port, ctx)(context) + + self.assertEqual(captured[0].fileno(), -1, "proxy socket was left open") + + def test_network_error_from_callback_propagates(self): + def callback(context): + raise OSError("proxy unreachable") + + # Not a ConfigurationError, so kms_request retries it. + with self.assertRaises(OSError): + _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + + @unittest.skipUnless(_HAVE_PYMONGOCRYPT, "pymongocrypt is not installed") + def test_client_encryption_accepts_callback(self): + def callback(context): + raise AssertionError("not called") + + client = self.simple_client() + encryption = ClientEncryption( + {"local": {"key": b"\x00" * 96}}, + "keyvault.datakeys", + client, + OPTS, + kms_connect_callback=callback, + ) + self.addCleanup(encryption.close) + self.assertIs(encryption._io_callbacks.opts._kms_connect_callback, callback) + + @unittest.skipUnless(_HAVE_PYMONGOCRYPT, "pymongocrypt is not installed") + def test_client_encryption_rejects_non_callable(self): + client = self.simple_client() + with self.assertRaisesRegex(TypeError, "kms_connect_callback must be callable"): + ClientEncryption( + {"local": {"key": b"\x00" * 96}}, + "keyvault.datakeys", + client, + OPTS, + kms_connect_callback="not-callable", # type: ignore[arg-type] + ) + class TestClientOptions(PyMongoTestCase): def test_default(self): @@ -252,9 +739,15 @@ def create_client_encryption( key_vault_client: MongoClient, codec_options: CodecOptions, kms_tls_options: Optional[Mapping[str, Any]] = None, + kms_connect_callback: Optional[Any] = None, ): client_encryption = ClientEncryption( - kms_providers, key_vault_namespace, key_vault_client, codec_options, kms_tls_options + kms_providers, + key_vault_namespace, + key_vault_client, + codec_options, + kms_tls_options, + kms_connect_callback=kms_connect_callback, ) self.addCleanup(client_encryption.close) return client_encryption @@ -267,9 +760,15 @@ def unmanaged_create_client_encryption( key_vault_client: MongoClient, codec_options: CodecOptions, kms_tls_options: Optional[Mapping[str, Any]] = None, + kms_connect_callback: Optional[Any] = None, ): client_encryption = ClientEncryption( - kms_providers, key_vault_namespace, key_vault_client, codec_options, kms_tls_options + kms_providers, + key_vault_namespace, + key_vault_client, + codec_options, + kms_tls_options, + kms_connect_callback=kms_connect_callback, ) return client_encryption @@ -1908,6 +2407,188 @@ def test_invalid_hostname_in_kms_certificate(self): self.client_encrypted.create_data_key("aws", master_key=key) +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", +} + + +# https://github.com/mongodb/specifications/blob/master/source/client-side-encryption/tests/README.md#kms-connect-callback +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(KMS_PROXY_HOST, KMS_PROXY_PORT)(context) + + def tls_callback(self, context): + self.callback_calls.append(context) + ctx = ssl.create_default_context(cafile=CA_PEM) + ctx.check_hostname = False + callback = HTTPProxyKMSConnect(KMS_PROXY_HOST, KMS_TLS_PROXY_PORT, ctx) + 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: + ctx = ssl.create_default_context(cafile=CA_PEM) + ctx.check_hostname = False + conn = http.client.HTTPSConnection( + f"{KMS_PROXY_HOST}:{KMS_TLS_PROXY_PORT}", context=ctx + ) + 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) + + self.assertGreaterEqual(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(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) + + # https://github.com/mongodb/specifications/blob/master/source/client-side-encryption/tests/README.md#kms-tls-options-tests class TestKmsTLSOptions(EncryptionIntegrationTest): @unittest.skipUnless(any(AWS_CREDS.values()), "AWS environment credentials are not set") diff --git a/test/test_pooling.py b/test/test_pooling.py index 64146a0e13..55f450d57a 100644 --- a/test/test_pooling.py +++ b/test/test_pooling.py @@ -34,13 +34,19 @@ from pymongo.hello import HelloCompat from pymongo.lock import _create_lock from pymongo.monitoring import _EventListeners +from pymongo.pool_shared import _wrap_socket_tls from test.utils import flaky, get_pool, joinall sys.path[0:0] = [""] from pymongo.socket_checker import SocketChecker from pymongo.synchronous.pool import Pool, PoolOptions -from test import IntegrationTest, client_context, unittest +from test import ( + IntegrationTest, + PyMongoTestCase, + client_context, + unittest, +) from test.helpers import ConcurrentRunner from test.utils_shared import CMAPListener, delay @@ -757,5 +763,18 @@ def test_certificate_error_is_not_labeled_overloaded(self): self.assertFalse(err.has_error_label("SystemOverloadedError")) +class TestWrapSocketTLS(PyMongoTestCase): + def test_wrap_socket_tls_without_ssl_context_returns_same_socket(self): + options = PoolOptions(socket_timeout=7.5) + left, right = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + + result = _wrap_socket_tls(left, ("kms.example.com", 443), options) + + self.assertIs(result, left) + self.assertEqual(result.gettimeout(), 7.5) + + if __name__ == "__main__": unittest.main() diff --git a/tools/synchro.py b/tools/synchro.py index bebf92c005..04a433763d 100644 --- a/tools/synchro.py +++ b/tools/synchro.py @@ -72,6 +72,8 @@ "_a_grid_out_property": "_grid_out_property", "AsyncClientEncryption": "ClientEncryption", "AsyncMongoCryptCallback": "MongoCryptCallback", + "AsyncKMSConnectCallback": "KMSConnectCallback", + "AsyncHTTPProxyKMSConnect": "HTTPProxyKMSConnect", "AsyncExplicitEncrypter": "ExplicitEncrypter", "AsyncAutoEncrypter": "AutoEncrypter", "AsyncContextManager": "ContextManager", @@ -127,6 +129,7 @@ "AsyncNetworkingInterface": "NetworkingInterface", "_configured_protocol_interface": "_configured_socket_interface", "_async_configured_socket": "_configured_socket", + "_async_wrap_socket_tls": "_wrap_socket_tls", "SpecRunnerTask": "SpecRunnerThread", "AsyncMockConnection": "MockConnection", "AsyncMockPool": "MockPool",