diff --git a/doc/changelog.rst b/doc/changelog.rst index cb3ded580a..9ea2f05c2e 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -16,6 +16,14 @@ 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 + :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. - 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/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 9ba2758f78..4ae63716e3 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -17,8 +17,11 @@ from __future__ import annotations import asyncio +import contextlib import functools +import inspect import socket +import ssl import time as time # noqa: PLC0414 # needed in sync version import uuid import weakref @@ -60,7 +63,9 @@ from pymongo.common import CONNECT_TIMEOUT from pymongo.daemon import _spawn_daemon from pymongo.encryption_options import ( + AsyncKMSConnectCallback, AutoEncryptionOpts, + KMSConnectContext, RangeOpts, StringOpts, # Re-exported for backwards compatibility: TextOpts is deprecated but must @@ -90,6 +95,8 @@ from pymongo.pool_options import PoolOptions from pymongo.pool_shared import ( _async_configured_socket, + _async_wrap_socket_tls, + _close_late_socket, _raise_connection_failure, ) from pymongo.read_concern import ReadConcern @@ -120,11 +127,108 @@ _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 callback return value that failed validation, best effort. + + ``_connect_kms`` raises on a contract violation instead of returning the + value, so no caller ever takes ownership of it. Close it here, tolerating + non-socket values and ``close()`` failures. + """ + 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]: + """Connect to a KMS host and perform the TLS handshake over the socket. + + Uses ``kms_connect_callback`` when one is provided, otherwise connects + directly, and always verifies against ``address`` (the KMS host). + """ + 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. + # A plain callable would block the event loop before we could reject it, + # so check the callback first. + if not _IS_SYNC: + callback_any: Any = kms_connect_callback + is_coro = inspect.iscoroutinefunction(callback_any) + if not is_coro and callable(callback_any): + is_coro = inspect.iscoroutinefunction(callback_any.__call__) + if not is_coro: + raise ConfigurationError( + "kms_connect_callback must be a coroutine function for the async API." + ) + # Any: the sync version's callback returns a plain socket. + result: Any = kms_connect_callback( + KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) + ) + remaining = _csot.remaining() + if remaining is None or _IS_SYNC: + # The synchronous API cannot interrupt a callback that has started + # running; honoring the deadline is the callback's contract there. + sock = await result + else: + # CSOT is cooperative: a callback that ignores the timeout could block + # past the deadline. Shield the task so stopping the wait does not + # cancel it mid-flight, and close any socket it yields later. + task = asyncio.ensure_future(result) + try: + sock = await asyncio.wait_for(asyncio.shield(task), remaining) + except asyncio.CancelledError: + task.add_done_callback(_close_late_socket) + raise + except asyncio.TimeoutError: + task.add_done_callback(_close_late_socket) + _raise_connection_failure( + address, + socket.timeout("timed out"), + timeout_details=_get_timeout_details(opts), + ) + 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)}." + ) + # wrap_socket refuses a non-blocking socket, so normalize the mode here. + try: + 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." + ) + # The callback may have consumed much of the CSOT budget, and wrapping + # resets the socket timeout, so recompute the remaining time here and for + # the KMS request that follows. + sock.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) try: - return await _async_configured_socket(address, opts) + conn = await _async_wrap_socket_tls(sock, address, opts) + except asyncio.CancelledError: + # The executor may still be wrapping the socket; close it so a TLS + # proxy's relay threads wind down instead of leaking. + sock.close() + raise except Exception as exc: _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) + conn.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) + return conn class _EncryptionIO(AsyncMongoCryptCallback): # type: ignore[misc] @@ -179,13 +283,6 @@ async def kms_request(self, kms_context: MongoCryptKmsContext) -> None: False, # disable_ocsp_endpoint_check _IS_SYNC, ) - # CSOT: set timeout for socket creation. - 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) if address[0].endswith(".sock"): raise ConfigurationError(f"Invalid KMS endpoint {endpoint!r}") @@ -193,8 +290,21 @@ async def kms_request(self, kms_context: MongoCryptKmsContext) -> None: if sleep_u: sleep_sec = float(sleep_u) / 1e6 await asyncio.sleep(sleep_sec) + # Set the connect timeout after the retry backoff so the budget + # reflects the sleep. + connect_timeout = max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001) + opts = PoolOptions( + connect_timeout=connect_timeout, + socket_timeout=connect_timeout, + ssl_context=ctx, + ) 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: @@ -230,6 +340,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): @@ -526,6 +638,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. @@ -595,7 +708,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, e.g. to route KMS requests through a proxy. It receives + a :class:`~pymongo.encryption_options.KMSConnectContext` and + returns a connected, unwrapped :class:`socket.socket`, over which + the driver performs the KMS TLS handshake. + Must be a coroutine function for the asynchronous API. + The callback is responsible for honoring ``context.timeout``. + Defaults to ``None``, meaning the driver connects to KMS hosts + directly. + + .. versionchanged:: 4.19 + Added the `kms_connect_callback` parameter. .. versionchanged:: 4.12 Added the `key_expiration_ms` parameter. .. versionchanged:: 4.0 @@ -639,6 +763,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 0d90fe3b19..b1d13700da 100644 --- a/pymongo/encryption_options.py +++ b/pymongo/encryption_options.py @@ -19,9 +19,11 @@ from __future__ import annotations +import socket import warnings -from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Optional, TypedDict +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 @@ -55,6 +57,38 @@ 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, + verifying against ``host`` rather than the peer actually reached. + + :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 + connect timeout, capped by the remaining time of an active operation + timeout (``timeoutMS``). + + .. 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.19 + """ + + 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] + + class AutoEncryptionOpts: """Options to configure automatic client-side field level encryption.""" @@ -75,6 +109,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. @@ -212,7 +247,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:`KMSConnectContext` and returns a connected, + unwrapped :class:`socket.socket`, over which the driver performs + the KMS TLS handshake. Must be a coroutine function for + :class:`~pymongo.asynchronous.mongo_client.AsyncMongoClient` and a + regular function for + :class:`~pymongo.synchronous.mongo_client.MongoClient`. Defaults + to ``None``, meaning the driver connects to KMS hosts directly. + + .. versionchanged:: 4.19 + Added the `kms_connect_callback` parameter. .. versionchanged:: 4.12 Added the `key_expiration_ms` parameter. .. versionchanged:: 4.2 @@ -259,6 +305,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 8cd546bda6..04663f6fc9 100644 --- a/pymongo/pool_shared.py +++ b/pymongo/pool_shared.py @@ -18,6 +18,7 @@ import asyncio import collections +import contextlib import functools import socket import ssl @@ -304,16 +305,29 @@ async def _async_create_connection(address: _Address, options: PoolOptions) -> s raise OSError("getaddrinfo failed") -async def _async_configured_socket( - address: _Address, options: PoolOptions +def _close_late_socket(future: asyncio.Future[Any]) -> None: + """Close a socket produced after its awaiting task was cancelled.""" + if not future.cancelled() and future.exception() is None: + # The callback may have returned a non-socket; close best effort. + close = getattr(future.result(), "close", None) + if callable(close): + with contextlib.suppress(Exception): + close() + + +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, e.g. + 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: @@ -326,13 +340,19 @@ async def _async_configured_socket( # to use SSLContext.check_hostname. if _has_sni(False): loop = asyncio.get_running_loop() - ssl_sock = await loop.run_in_executor( - None, - functools.partial(ssl_context.wrap_socket, sock, server_hostname=host), # type: ignore[assignment, misc, unused-ignore] - ) + wrap = functools.partial(ssl_context.wrap_socket, sock, server_hostname=host) # type: ignore[assignment, misc, unused-ignore] else: loop = asyncio.get_running_loop() - ssl_sock = await loop.run_in_executor(None, ssl_context.wrap_socket, sock) # type: ignore[assignment, misc, unused-ignore] + wrap = functools.partial(ssl_context.wrap_socket, sock) # type: ignore[assignment, misc, unused-ignore] + # Shield the executor future: wrap_socket hands the fd to a new + # SSLSocket, so cancellation must not orphan the result it produces. + future = loop.run_in_executor(None, wrap) + try: + ssl_sock = await asyncio.shield(future) + except asyncio.CancelledError: + future.add_done_callback(_close_late_socket) + sock.close() + raise except _CertificateError: sock.close() # Raise _CertificateError directly like we do after match_hostname @@ -360,6 +380,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, @@ -510,14 +543,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, e.g. + 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: @@ -559,6 +597,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 e7d8a366ca..ac2f710382 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -16,8 +16,12 @@ from __future__ import annotations +import asyncio +import contextlib import functools +import inspect import socket +import ssl import time as time # noqa: PLC0414 # needed in sync version import uuid import weakref @@ -56,6 +60,8 @@ from pymongo.daemon import _spawn_daemon from pymongo.encryption_options import ( AutoEncryptionOpts, + KMSConnectCallback, + KMSConnectContext, RangeOpts, StringOpts, # Re-exported for backwards compatibility: TextOpts is deprecated but must @@ -84,8 +90,10 @@ from pymongo.operations import UpdateOne from pymongo.pool_options import PoolOptions from pymongo.pool_shared import ( + _close_late_socket, _configured_socket, _raise_connection_failure, + _wrap_socket_tls, ) from pymongo.read_concern import ReadConcern from pymongo.results import DeleteResult @@ -119,11 +127,108 @@ _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 callback return value that failed validation, best effort. + + ``_connect_kms`` raises on a contract violation instead of returning the + value, so no caller ever takes ownership of it. Close it here, tolerating + non-socket values and ``close()`` failures. + """ + 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]: + """Connect to a KMS host and perform the TLS handshake over the socket. + + Uses ``kms_connect_callback`` when one is provided, otherwise connects + directly, and always verifies against ``address`` (the KMS host). + """ + 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. + # A plain callable would block the event loop before we could reject it, + # so check the callback first. + if not _IS_SYNC: + callback_any: Any = kms_connect_callback + is_coro = inspect.iscoroutinefunction(callback_any) + if not is_coro and callable(callback_any): + is_coro = inspect.iscoroutinefunction(callback_any.__call__) + if not is_coro: + raise ConfigurationError( + "kms_connect_callback must be a coroutine function for the async API." + ) + # Any: the sync version's callback returns a plain socket. + result: Any = kms_connect_callback( + KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) + ) + remaining = _csot.remaining() + if remaining is None or _IS_SYNC: + # The synchronous API cannot interrupt a callback that has started + # running; honoring the deadline is the callback's contract there. + sock = result + else: + # CSOT is cooperative: a callback that ignores the timeout could block + # past the deadline. Shield the task so stopping the wait does not + # cancel it mid-flight, and close any socket it yields later. + task = asyncio.ensure_future(result) + try: + sock = asyncio.wait_for(asyncio.shield(task), remaining) + except asyncio.CancelledError: + task.add_done_callback(_close_late_socket) + raise + except asyncio.TimeoutError: + task.add_done_callback(_close_late_socket) + _raise_connection_failure( + address, + socket.timeout("timed out"), + timeout_details=_get_timeout_details(opts), + ) + 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)}." + ) + # wrap_socket refuses a non-blocking socket, so normalize the mode here. + try: + 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." + ) + # The callback may have consumed much of the CSOT budget, and wrapping + # resets the socket timeout, so recompute the remaining time here and for + # the KMS request that follows. + sock.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) try: - return _configured_socket(address, opts) + conn = _wrap_socket_tls(sock, address, opts) + except asyncio.CancelledError: + # The executor may still be wrapping the socket; close it so a TLS + # proxy's relay threads wind down instead of leaking. + sock.close() + raise except Exception as exc: _raise_connection_failure(address, exc, timeout_details=_get_timeout_details(opts)) + conn.settimeout(max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001)) + return conn class _EncryptionIO(MongoCryptCallback): # type: ignore[misc] @@ -178,13 +283,6 @@ def kms_request(self, kms_context: MongoCryptKmsContext) -> None: False, # disable_ocsp_endpoint_check _IS_SYNC, ) - # CSOT: set timeout for socket creation. - 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) if address[0].endswith(".sock"): raise ConfigurationError(f"Invalid KMS endpoint {endpoint!r}") @@ -192,8 +290,21 @@ def kms_request(self, kms_context: MongoCryptKmsContext) -> None: if sleep_u: sleep_sec = float(sleep_u) / 1e6 time.sleep(sleep_sec) + # Set the connect timeout after the retry backoff so the budget + # reflects the sleep. + connect_timeout = max(_csot.clamp_remaining(_KMS_CONNECT_TIMEOUT), 0.001) + opts = PoolOptions( + connect_timeout=connect_timeout, + socket_timeout=connect_timeout, + ssl_context=ctx, + ) 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: @@ -229,6 +340,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): @@ -523,6 +636,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. @@ -592,7 +706,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, e.g. to route KMS requests through a proxy. It receives + a :class:`~pymongo.encryption_options.KMSConnectContext` and + returns a connected, unwrapped :class:`socket.socket`, over which + the driver performs the KMS TLS handshake. + Must be a regular function. + The callback is responsible for honoring ``context.timeout``. + Defaults to ``None``, meaning the driver connects to KMS hosts + directly. + + .. versionchanged:: 4.19 + Added the `kms_connect_callback` parameter. .. versionchanged:: 4.12 Added the `key_expiration_ms` parameter. .. versionchanged:: 4.0 @@ -632,6 +757,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 128f26feb4..e22d185f77 100644 --- a/test/asynchronous/test_encryption.py +++ b/test/asynchronous/test_encryption.py @@ -60,7 +60,11 @@ 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, +) from pymongo.asynchronous.helpers import anext from pymongo.asynchronous.mongo_client import AsyncMongoClient from pymongo.cursor_shared import CursorType @@ -222,6 +226,9 @@ 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. + + class TestClientOptions(AsyncPyMongoTestCase): async def test_default(self): client = self.simple_client(connect=False) @@ -316,9 +323,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 @@ -331,9 +344,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 @@ -1988,6 +2007,9 @@ 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. + + # 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_kms_connect.py b/test/asynchronous/test_kms_connect.py new file mode 100644 index 0000000000..6a577b90ee --- /dev/null +++ b/test/asynchronous/test_kms_connect.py @@ -0,0 +1,449 @@ +"""Tests for the KMS connect callback.""" + +from __future__ import annotations + +import asyncio +import dataclasses +import os +import socket +import ssl +import threading +import time +import unittest +from asyncio.trsock import TransportSocket +from unittest import mock + +import pytest + +import pymongo +from pymongo.asynchronous.encryption import ( + AsyncClientEncryption, + _connect_kms, + _EncryptionIO, + _wrap_encryption_errors, +) +from pymongo.encryption_options import ( + AutoEncryptionOpts, + KMSConnectContext, +) +from pymongo.errors import ConfigurationError, ConnectionFailure, EncryptionError, NetworkTimeout +from pymongo.pool_options import PoolOptions +from pymongo.ssl_support import get_ssl_context +from test.asynchronous import AsyncPyMongoTestCase +from test.asynchronous.test_encryption import OPTS +from test.helpers_shared import CA_PEM, CERT_PATH, CLIENT_PEM + +_IS_SYNC = False + +pytestmark = pytest.mark.encryption + +_KMS_ADDRESS = ("kms.example.com", 443) + + +def _tls_server_context(cert=CLIENT_PEM): + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + ctx.load_cert_chain(cert) + return ctx + + +def _client_tls_context(verify=False): + # verify=False matches the driver's test mode: the local certs don't verify. + if verify: + return get_ssl_context(None, None, CA_PEM, None, False, False, False, _IS_SYNC) + return get_ssl_context(None, None, None, None, True, True, False, _IS_SYNC) + + +async def _run_blocking(func, *args): + """Run a blocking callable off the event loop (inline in the sync version).""" + if _IS_SYNC: + return func(*args) + return await asyncio.get_running_loop().run_in_executor(None, func, *args) + + +def _callback_returning(value): + """A kms_connect_callback that always produces ``value``.""" + + async def callback(context): + return value + + return callback + + +class TestKmsConnectCallbackUnit(AsyncPyMongoTestCase): + """Contract checks for kms_connect_callback that need no KMS server.""" + + @staticmethod + def _pool_options(ssl_context=None): + return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=ssl_context) + + def _listen(self, backlog=1): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(backlog) + self.addCleanup(listener.close) + return listener + + def _socketpair(self): + left, right = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + return left, right + + 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] + + async def test_non_socket_return_raises_configuration_error(self): + with self.assertRaisesRegex(ConfigurationError, "must return a connected"): + await _connect_kms( + _KMS_ADDRESS, self._pool_options(), _callback_returning("not-a-socket"), 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 = self._socketpair() + # 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) + + with self.assertRaisesRegex(ConfigurationError, "unwrapped"): + await _connect_kms( + _KMS_ADDRESS, self._pool_options(), _callback_returning(wrapped), 10.0 + ) + + async def test_context_receives_host_port_and_timeout(self): + received = [] + left, _right = self._socketpair() + + 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_ADDRESS, 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 = _tls_server_context() + listener = self._listen() + + 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() + + options = self._pool_options(_client_tls_context()) + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + sock.setblocking(False) + return sock + + async def callback(context): + return await _run_blocking(connect) + + conn = await _connect_kms(listener.getsockname(), options, callback, 10.0) + self.addCleanup(conn.close) + self.assertIsNotNone(conn.gettimeout()) + + async def test_tls_verification_targets_the_kms_host(self): + # The handshake must verify against the KMS address, not the peer the + # callback connected to. The server cert covers 127.0.0.1 (the peer) + # and localhost, but not the KMS hostname used below, so only + # address-based verification produces this outcome. + server_ctx = _tls_server_context(os.path.join(CERT_PATH, "server.pem")) + listener = self._listen(2) + + def serve(): + for _ in range(2): + try: + conn, _ = listener.accept() + server_ctx.wrap_socket(conn, server_side=True).close() + except (OSError, ssl.SSLError): + # The mismatched-name attempt fails mid-handshake. + pass + + threading.Thread(target=serve, daemon=True).start() + + # Full verification: trusted CA, invalid certs and hostnames rejected. + options = self._pool_options(_client_tls_context(verify=True)) + + created = [] + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + created.append(sock) + return sock + + async def callback(context): + return await _run_blocking(connect) + + port = listener.getsockname()[1] + # The cert covers localhost: verifying against the KMS address succeeds. + conn = await _connect_kms(("localhost", port), options, callback, 10.0) + self.addCleanup(conn.close) + # TLS-wrapped in either SSL flavor: a new object, not the plain socket. + self.assertIsNot(conn, created[0]) + # The cert does not cover this name: verification must fail even + # though the peer (127.0.0.1) presents a cert valid for itself. + with self.assertRaises(ConnectionFailure): + await _connect_kms(("kms.example.com", port), options, callback, 10.0) + + async def test_asyncio_transport_socket_is_rejected(self): + # get_extra_info("socket") is a TransportSocket, not a socket.socket. + left, _right = self._socketpair() + + with self.assertRaisesRegex(ConfigurationError, "TransportSocket"): + await _connect_kms( + _KMS_ADDRESS, self._pool_options(), _callback_returning(TransportSocket(left)), 10.0 + ) + + async def test_cancelled_tls_wrap_closes_late_socket(self): + # A cancelled wrap can leave the executor producing an SSLSocket; the + # done callback must close it. + if _IS_SYNC: + raise unittest.SkipTest("the cancel-safe wrap is an async path") + from pymongo.pool_shared import _close_late_socket + + left, _right = self._socketpair() + future = asyncio.get_running_loop().create_future() + future.set_result(left) + self.assertNotEqual(left.fileno(), -1) + _close_late_socket(future) + self.assertEqual(left.fileno(), -1) + + async def test_non_coroutine_callback_is_rejected(self): + # A plain def must be rejected before it blocks the event loop. + if _IS_SYNC: + raise unittest.SkipTest("a regular function is correct for the sync API") + + entered = [] + + def callback(context): + entered.append(context) + return None + + with self.assertRaisesRegex(ConfigurationError, "coroutine function"): + await _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 10.0) + self.assertEqual(entered, [], "invalid callback must not be entered") + + async def test_unconnected_socket_from_callback_is_rejected(self): + # An unconnected socket would fail later as a transient error and be retried. + bare = socket.socket() + self.addCleanup(bare.close) + + with self.assertRaisesRegex(ConfigurationError, "already connected"): + await _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(bare), 10.0) + + async def test_datagram_socket_from_callback_is_rejected(self): + # TLS on a connected UDP socket 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()) + + with self.assertRaisesRegex(ConfigurationError, "stream socket"): + await _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(left), 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): + # 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_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_ADDRESS, self._pool_options(), callback, 10.0) + + async def test_csot_deadline_stops_a_hung_callback(self): + # A callback that ignores the timeout cannot block past the CSOT + # deadline, and a socket it yields later must be closed. + if _IS_SYNC: + raise unittest.SkipTest("the sync API cannot interrupt a callback") + + left, _right = self._socketpair() + + async def hung_callback(context): + await asyncio.sleep(0.5) + return left + + with self.assertRaises(NetworkTimeout): + with pymongo.timeout(0.1): + await _connect_kms(_KMS_ADDRESS, self._pool_options(), hung_callback, 10.0) + self.assertNotEqual(left.fileno(), -1) + # Let the shielded callback finish; the driver closes the late result. + await asyncio.sleep(0.75) + self.assertEqual(left.fileno(), -1) + + async def test_cancelling_kms_connect_closes_the_callback_socket(self): + # Cancelling during the TLS handshake must close the callback's socket, + # so a TLS proxy's relay threads wind down. + if _IS_SYNC: + raise unittest.SkipTest("cancellation is an async-only behavior") + + server_ctx = _tls_server_context() + listener = self._listen() + gate = threading.Event() + eof = threading.Event() + + def stub_server(): + conn = None + try: + conn, _ = listener.accept() + # The cancel may land before or after the executor starts the + # handshake. Peek for the ClientHello without consuming it, or + # for EOF if the driver closed it, before wrap_socket detaches conn. + while True: + data = conn.recv(4096, socket.MSG_PEEK) + if not data: + eof.set() + return + if data[:1] == b"\x16": # TLS handshake record + break + # Hold the handshake open until the test has cancelled. + if not gate.wait(5): + return + tls = server_ctx.wrap_socket(conn, server_side=True, do_handshake_on_connect=False) + try: + tls.do_handshake() + # A discarded connection may end in a reset rather than a + # clean EOF; either proves the driver closed it. + while tls.recv(4096): + pass + except OSError: + pass + eof.set() + tls.close() + except OSError: + pass + finally: + if conn is not None: + conn.close() + + threading.Thread(target=stub_server, daemon=True).start() + + options = self._pool_options(_client_tls_context()) + socks = [] + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + socks.append(sock) + return sock + + async def callback(context): + return await _run_blocking(connect) + + # The sync version returns a socket instead of a coroutine, so both + # error codes are needed depending on the version being checked. + pending = _connect_kms(listener.getsockname(), options, callback, 10.0) + task = asyncio.ensure_future(pending) # type: ignore[type-var,arg-type] + for _ in range(100): + if socks: + break + await asyncio.sleep(0.01) + self.assertTrue(socks, "callback was never invoked") + # Bias the cancel to land mid-handshake; the stub handles the earlier + # window too. + await asyncio.sleep(0.1) + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + # The late SSLSocket (or raw socket) must be closed; the stub sees EOF. + gate.set() + for _ in range(50): + if eof.is_set(): + break + await asyncio.sleep(0.1) + self.assertTrue(eof.is_set(), "driver never closed the callback socket") + + 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) + + 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] + ) diff --git a/test/asynchronous/test_pooling.py b/test/asynchronous/test_pooling.py index 25da5134b9..9bd221f6e7 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 @@ -919,5 +925,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 adb6005ea1..2f9ed382ca 100644 --- a/test/test_encryption.py +++ b/test/test_encryption.py @@ -82,7 +82,11 @@ ) from pymongo.operations import InsertOne, ReplaceOne, UpdateOne from pymongo.synchronous import encryption -from pymongo.synchronous.encryption import Algorithm, ClientEncryption, QueryType +from pymongo.synchronous.encryption import ( + Algorithm, + ClientEncryption, + QueryType, +) from pymongo.synchronous.helpers import next from pymongo.synchronous.mongo_client import MongoClient from pymongo.write_concern import WriteConcern @@ -222,6 +226,9 @@ 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. + + class TestClientOptions(PyMongoTestCase): def test_default(self): client = self.simple_client(connect=False) @@ -316,9 +323,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 @@ -331,9 +344,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 @@ -1980,6 +1999,9 @@ 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. + + # 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_kms_connect.py b/test/test_kms_connect.py new file mode 100644 index 0000000000..8bd59ca689 --- /dev/null +++ b/test/test_kms_connect.py @@ -0,0 +1,447 @@ +"""Tests for the KMS connect callback.""" + +from __future__ import annotations + +import asyncio +import dataclasses +import os +import socket +import ssl +import threading +import time +import unittest +from asyncio.trsock import TransportSocket +from unittest import mock + +import pytest + +import pymongo +from pymongo.encryption_options import ( + AutoEncryptionOpts, + KMSConnectContext, +) +from pymongo.errors import ConfigurationError, ConnectionFailure, EncryptionError, NetworkTimeout +from pymongo.pool_options import PoolOptions +from pymongo.ssl_support import get_ssl_context +from pymongo.synchronous.encryption import ( + ClientEncryption, + _connect_kms, + _EncryptionIO, + _wrap_encryption_errors, +) +from test import PyMongoTestCase +from test.helpers_shared import CA_PEM, CERT_PATH, CLIENT_PEM +from test.test_encryption import OPTS + +_IS_SYNC = True + +pytestmark = pytest.mark.encryption + +_KMS_ADDRESS = ("kms.example.com", 443) + + +def _tls_server_context(cert=CLIENT_PEM): + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + ctx.load_cert_chain(cert) + return ctx + + +def _client_tls_context(verify=False): + # verify=False matches the driver's test mode: the local certs don't verify. + if verify: + return get_ssl_context(None, None, CA_PEM, None, False, False, False, _IS_SYNC) + return get_ssl_context(None, None, None, None, True, True, False, _IS_SYNC) + + +def _run_blocking(func, *args): + """Run a blocking callable off the event loop (inline in the sync version).""" + if _IS_SYNC: + return func(*args) + return asyncio.get_running_loop().run_in_executor(None, func, *args) + + +def _callback_returning(value): + """A kms_connect_callback that always produces ``value``.""" + + def callback(context): + return value + + return callback + + +class TestKmsConnectCallbackUnit(PyMongoTestCase): + """Contract checks for kms_connect_callback that need no KMS server.""" + + @staticmethod + def _pool_options(ssl_context=None): + return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=ssl_context) + + def _listen(self, backlog=1): + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(backlog) + self.addCleanup(listener.close) + return listener + + def _socketpair(self): + left, right = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + return left, right + + 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] + + def test_non_socket_return_raises_configuration_error(self): + with self.assertRaisesRegex(ConfigurationError, "must return a connected"): + _connect_kms( + _KMS_ADDRESS, self._pool_options(), _callback_returning("not-a-socket"), 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 = self._socketpair() + # 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) + + with self.assertRaisesRegex(ConfigurationError, "unwrapped"): + _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(wrapped), 10.0) + + def test_context_receives_host_port_and_timeout(self): + received = [] + left, _right = self._socketpair() + + 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_ADDRESS, 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 = _tls_server_context() + listener = self._listen() + + 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() + + options = self._pool_options(_client_tls_context()) + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + sock.setblocking(False) + return sock + + def callback(context): + return _run_blocking(connect) + + conn = _connect_kms(listener.getsockname(), options, callback, 10.0) + self.addCleanup(conn.close) + self.assertIsNotNone(conn.gettimeout()) + + def test_tls_verification_targets_the_kms_host(self): + # The handshake must verify against the KMS address, not the peer the + # callback connected to. The server cert covers 127.0.0.1 (the peer) + # and localhost, but not the KMS hostname used below, so only + # address-based verification produces this outcome. + server_ctx = _tls_server_context(os.path.join(CERT_PATH, "server.pem")) + listener = self._listen(2) + + def serve(): + for _ in range(2): + try: + conn, _ = listener.accept() + server_ctx.wrap_socket(conn, server_side=True).close() + except (OSError, ssl.SSLError): + # The mismatched-name attempt fails mid-handshake. + pass + + threading.Thread(target=serve, daemon=True).start() + + # Full verification: trusted CA, invalid certs and hostnames rejected. + options = self._pool_options(_client_tls_context(verify=True)) + + created = [] + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + created.append(sock) + return sock + + def callback(context): + return _run_blocking(connect) + + port = listener.getsockname()[1] + # The cert covers localhost: verifying against the KMS address succeeds. + conn = _connect_kms(("localhost", port), options, callback, 10.0) + self.addCleanup(conn.close) + # TLS-wrapped in either SSL flavor: a new object, not the plain socket. + self.assertIsNot(conn, created[0]) + # The cert does not cover this name: verification must fail even + # though the peer (127.0.0.1) presents a cert valid for itself. + with self.assertRaises(ConnectionFailure): + _connect_kms(("kms.example.com", port), options, callback, 10.0) + + def test_asyncio_transport_socket_is_rejected(self): + # get_extra_info("socket") is a TransportSocket, not a socket.socket. + left, _right = self._socketpair() + + with self.assertRaisesRegex(ConfigurationError, "TransportSocket"): + _connect_kms( + _KMS_ADDRESS, self._pool_options(), _callback_returning(TransportSocket(left)), 10.0 + ) + + def test_cancelled_tls_wrap_closes_late_socket(self): + # A cancelled wrap can leave the executor producing an SSLSocket; the + # done callback must close it. + if _IS_SYNC: + raise unittest.SkipTest("the cancel-safe wrap is an async path") + from pymongo.pool_shared import _close_late_socket + + left, _right = self._socketpair() + future = asyncio.get_running_loop().create_future() + future.set_result(left) + self.assertNotEqual(left.fileno(), -1) + _close_late_socket(future) + self.assertEqual(left.fileno(), -1) + + def test_non_coroutine_callback_is_rejected(self): + # A plain def must be rejected before it blocks the event loop. + if _IS_SYNC: + raise unittest.SkipTest("a regular function is correct for the sync API") + + entered = [] + + def callback(context): + entered.append(context) + return None + + with self.assertRaisesRegex(ConfigurationError, "coroutine function"): + _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 10.0) + self.assertEqual(entered, [], "invalid callback must not be entered") + + def test_unconnected_socket_from_callback_is_rejected(self): + # An unconnected socket would fail later as a transient error and be retried. + bare = socket.socket() + self.addCleanup(bare.close) + + with self.assertRaisesRegex(ConfigurationError, "already connected"): + _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(bare), 10.0) + + def test_datagram_socket_from_callback_is_rejected(self): + # TLS on a connected UDP socket 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()) + + with self.assertRaisesRegex(ConfigurationError, "stream socket"): + _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(left), 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): + # 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_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_ADDRESS, self._pool_options(), callback, 10.0) + + def test_csot_deadline_stops_a_hung_callback(self): + # A callback that ignores the timeout cannot block past the CSOT + # deadline, and a socket it yields later must be closed. + if _IS_SYNC: + raise unittest.SkipTest("the sync API cannot interrupt a callback") + + left, _right = self._socketpair() + + def hung_callback(context): + time.sleep(0.5) + return left + + with self.assertRaises(NetworkTimeout): + with pymongo.timeout(0.1): + _connect_kms(_KMS_ADDRESS, self._pool_options(), hung_callback, 10.0) + self.assertNotEqual(left.fileno(), -1) + # Let the shielded callback finish; the driver closes the late result. + time.sleep(0.75) + self.assertEqual(left.fileno(), -1) + + def test_cancelling_kms_connect_closes_the_callback_socket(self): + # Cancelling during the TLS handshake must close the callback's socket, + # so a TLS proxy's relay threads wind down. + if _IS_SYNC: + raise unittest.SkipTest("cancellation is an async-only behavior") + + server_ctx = _tls_server_context() + listener = self._listen() + gate = threading.Event() + eof = threading.Event() + + def stub_server(): + conn = None + try: + conn, _ = listener.accept() + # The cancel may land before or after the executor starts the + # handshake. Peek for the ClientHello without consuming it, or + # for EOF if the driver closed it, before wrap_socket detaches conn. + while True: + data = conn.recv(4096, socket.MSG_PEEK) + if not data: + eof.set() + return + if data[:1] == b"\x16": # TLS handshake record + break + # Hold the handshake open until the test has cancelled. + if not gate.wait(5): + return + tls = server_ctx.wrap_socket(conn, server_side=True, do_handshake_on_connect=False) + try: + tls.do_handshake() + # A discarded connection may end in a reset rather than a + # clean EOF; either proves the driver closed it. + while tls.recv(4096): + pass + except OSError: + pass + eof.set() + tls.close() + except OSError: + pass + finally: + if conn is not None: + conn.close() + + threading.Thread(target=stub_server, daemon=True).start() + + options = self._pool_options(_client_tls_context()) + socks = [] + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + socks.append(sock) + return sock + + def callback(context): + return _run_blocking(connect) + + # The sync version returns a socket instead of a coroutine, so both + # error codes are needed depending on the version being checked. + pending = _connect_kms(listener.getsockname(), options, callback, 10.0) + task = asyncio.ensure_future(pending) # type: ignore[type-var,arg-type] + for _ in range(100): + if socks: + break + time.sleep(0.01) + self.assertTrue(socks, "callback was never invoked") + # Bias the cancel to land mid-handshake; the stub handles the earlier + # window too. + time.sleep(0.1) + task.cancel() + with self.assertRaises(asyncio.CancelledError): + task + # The late SSLSocket (or raw socket) must be closed; the stub sees EOF. + gate.set() + for _ in range(50): + if eof.is_set(): + break + time.sleep(0.1) + self.assertTrue(eof.is_set(), "driver never closed the callback socket") + + 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) + + 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] + ) diff --git a/test/test_pooling.py b/test/test_pooling.py index 3a81bc796f..d690de0ffb 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 @@ -917,5 +923,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..b5c6b2873a 100644 --- a/tools/synchro.py +++ b/tools/synchro.py @@ -72,6 +72,7 @@ "_a_grid_out_property": "_grid_out_property", "AsyncClientEncryption": "ClientEncryption", "AsyncMongoCryptCallback": "MongoCryptCallback", + "AsyncKMSConnectCallback": "KMSConnectCallback", "AsyncExplicitEncrypter": "ExplicitEncrypter", "AsyncAutoEncrypter": "AutoEncrypter", "AsyncContextManager": "ContextManager", @@ -127,6 +128,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", @@ -297,6 +299,13 @@ def translate_docstrings(lines: list[str]) -> list[str]: lines[i] = lines[i].replace("an asynchronous", "a") if "An asynchronous" in lines[i]: lines[i] = lines[i].replace("An asynchronous", "A") + # This sentence states the callback contract, whose meaning + # would invert under the async -> sync word replacements. + if "Must be a coroutine function for the asynchronous API." in lines[i]: + lines[i] = lines[i].replace( + "Must be a coroutine function for the asynchronous API.", + "Must be a regular function.", + ) # This ensures docstring links are for `pymongo.X` instead of `pymongo.synchronous.X` if "pymongo.asynchronous" in lines[i] and "import" not in lines[i]: lines[i] = lines[i].replace("pymongo.asynchronous", "pymongo")