From 6c92519559f1ff360debc7583873cc0bd12e4cf8 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 06:59:11 -0500 Subject: [PATCH 1/9] PYTHON-5805 Add kms_connect_callback for KMS connections Add a kms_connect_callback option to AutoEncryptionOpts, ClientEncryption, and AsyncClientEncryption. The callable receives a KMSConnectContext (host, port, timeout) and returns a connected, unwrapped socket over which the driver performs the KMS TLS handshake, so verification still targets the KMS host rather than the peer reached. This enables routing KMS requests through custom channels such as proxies. Split TLS wrapping out of the configured socket helpers so the callback path reuses it, and shield the async wrap so a cancelled or timed-out wait still closes a late-produced socket. --- doc/changelog.rst | 8 + pymongo/asynchronous/encryption.py | 148 +++++++++- pymongo/encryption_options.py | 56 +++- pymongo/pool_shared.py | 77 ++++- pymongo/synchronous/encryption.py | 149 +++++++++- test/asynchronous/test_encryption.py | 28 +- test/asynchronous/test_kms_connect.py | 393 ++++++++++++++++++++++++++ test/asynchronous/test_pooling.py | 21 +- test/test_encryption.py | 28 +- test/test_kms_connect.py | 391 +++++++++++++++++++++++++ test/test_pooling.py | 21 +- tools/synchro.py | 14 + 12 files changed, 1287 insertions(+), 47 deletions(-) create mode 100644 test/asynchronous/test_kms_connect.py create mode 100644 test/test_kms_connect.py 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..0eea43d898 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 rejected kms_connect_callback return value, best effort. + + Nothing else will close it: _connect_kms raises before the result reaches + the caller's ``finally``. + """ + 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." + ) + # Typed as Any so the generated synchronous flavor type-checks: the sync + # callback returns a plain socket, which is not awaitable. + 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,19 @@ 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`, over + which the driver performs the KMS TLS handshake. The callback + must be a coroutine function for the asynchronous API; a plain callable is rejected before it can block the event loop. + When a CSOT timeout is active, the driver stops waiting at the + deadline and closes any socket the callback yields later. + 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 +764,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..f77e461704 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,37 @@ 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 left in the timeout budget, or the default KMS + connect timeout when no timeout is active. + + .. 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 +108,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 +246,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 +304,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..b89ad933f1 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 rejected kms_connect_callback return value, best effort. + + Nothing else will close it: _connect_kms raises before the result reaches + the caller's ``finally``. + """ + 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." + ) + # Typed as Any so the generated synchronous flavor type-checks: the sync + # callback returns a plain socket, which is not awaitable. + 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,19 @@ 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`, over + which the driver performs the KMS TLS handshake. The callback + must be a regular function; the async API requires a coroutine function and rejects plain callables before they can block the event loop. + When a CSOT timeout is active, the driver stops waiting at the + deadline and closes any socket the callback yields later. + 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 +758,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..dc5d0e31e0 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 unit and prose 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 unit and prose 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..994ec5ca21 --- /dev/null +++ b/test/asynchronous/test_kms_connect.py @@ -0,0 +1,393 @@ +"""Tests for the KMS connect callback.""" + +from __future__ import annotations + +import asyncio +import dataclasses +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, 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 CLIENT_PEM + +_IS_SYNC = False + +pytestmark = pytest.mark.encryption + + +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_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): + 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_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 = socket.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) + self.addCleanup(right.close) + + 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.example.com", 443), 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) + + 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_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()) + + 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): + # 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.example.com", 443), 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 = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + + 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.example.com", 443), 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 = 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) + 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() + + 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) + socks = [] + + async def callback(context): + sock = await asyncio.get_running_loop().run_in_executor( + None, + lambda: socket.create_connection(listener.getsockname(), timeout=10), + ) + socks.append(sock) + return sock + + # The sync flavor returns a socket instead of a coroutine, so both + # error codes are needed depending on the flavor being checked. + connect = _connect_kms(listener.getsockname(), options, callback, 10.0) + task = asyncio.ensure_future(connect) # 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..01cf8cfc83 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 unit and prose 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 unit and prose 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..25c96191ab --- /dev/null +++ b/test/test_kms_connect.py @@ -0,0 +1,391 @@ +"""Tests for the KMS connect callback.""" + +from __future__ import annotations + +import asyncio +import dataclasses +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, 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 CLIENT_PEM +from test.test_encryption import OPTS + +_IS_SYNC = True + +pytestmark = pytest.mark.encryption + + +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_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): + 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_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 = socket.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) + self.addCleanup(right.close) + + 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.example.com", 443), 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) + + 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_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()) + + 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): + # 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.example.com", 443), 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 = socket.socketpair() + self.addCleanup(left.close) + self.addCleanup(right.close) + + def hung_callback(context): + time.sleep(0.5) + return left + + with self.assertRaises(NetworkTimeout): + with pymongo.timeout(0.1): + _connect_kms(("kms.example.com", 443), 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 = 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) + 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() + + 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) + socks = [] + + def callback(context): + sock = asyncio.get_running_loop().run_in_executor( + None, + lambda: socket.create_connection(listener.getsockname(), timeout=10), + ) + socks.append(sock) + return sock + + # The sync flavor returns a socket instead of a coroutine, so both + # error codes are needed depending on the flavor being checked. + connect = _connect_kms(listener.getsockname(), options, callback, 10.0) + task = asyncio.ensure_future(connect) # 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..bd5edebda3 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,18 @@ 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; a plain callable is rejected before it can block the event loop" + in lines[i] + ): + lines[i] = lines[i].replace( + "must be a coroutine function for the asynchronous API; a plain callable is rejected before it can block the event loop", + "must be a regular function; the async API requires a " + "coroutine function and rejects plain callables before " + "they can block the event loop", + ) # 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") From fd1b5a076396c5fa8955a5f6daa33fd0f08291d2 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 09:09:36 -0500 Subject: [PATCH 2/9] PYTHON-5805 Address follow-up review - Add a regression test that the KMS TLS handshake verifies against the KMS address, not the callback's connected peer. - Drop the stale prose-test references in test_encryption.py comments; prose tests moved with the PYTHON-6147 split. - Correct the KMSConnectContext.timeout docstring: it is the default KMS connect timeout capped by any active operation timeout. --- pymongo/encryption_options.py | 5 ++- test/asynchronous/test_encryption.py | 4 +- test/asynchronous/test_kms_connect.py | 55 ++++++++++++++++++++++++++- test/test_encryption.py | 4 +- test/test_kms_connect.py | 55 ++++++++++++++++++++++++++- 5 files changed, 113 insertions(+), 10 deletions(-) diff --git a/pymongo/encryption_options.py b/pymongo/encryption_options.py index f77e461704..b1d13700da 100644 --- a/pymongo/encryption_options.py +++ b/pymongo/encryption_options.py @@ -67,8 +67,9 @@ class KMSConnectContext: :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, or the default KMS - connect timeout when no timeout is active. + :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 diff --git a/test/asynchronous/test_encryption.py b/test/asynchronous/test_encryption.py index dc5d0e31e0..e22d185f77 100644 --- a/test/asynchronous/test_encryption.py +++ b/test/asynchronous/test_encryption.py @@ -226,7 +226,7 @@ async def test_init_kms_tls_options(self): self.assertEqual(ctx.verify_mode, ssl.CERT_REQUIRED) -# KMS connect callback unit and prose tests live in test_kms_connect.py. +# KMS connect callback tests live in test_kms_connect.py. class TestClientOptions(AsyncPyMongoTestCase): @@ -2007,7 +2007,7 @@ async def test_invalid_hostname_in_kms_certificate(self): await self.client_encrypted.create_data_key("aws", master_key=key) -# KMS connect callback unit and prose tests live in test_kms_connect.py. +# 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 diff --git a/test/asynchronous/test_kms_connect.py b/test/asynchronous/test_kms_connect.py index 994ec5ca21..99258cadda 100644 --- a/test/asynchronous/test_kms_connect.py +++ b/test/asynchronous/test_kms_connect.py @@ -4,6 +4,7 @@ import asyncio import dataclasses +import os import socket import ssl import threading @@ -25,12 +26,12 @@ AutoEncryptionOpts, KMSConnectContext, ) -from pymongo.errors import ConfigurationError, EncryptionError, NetworkTimeout +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 CLIENT_PEM +from test.helpers_shared import CA_PEM, CERT_PATH, CLIENT_PEM _IS_SYNC = False @@ -144,6 +145,56 @@ async def callback(context): 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 = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_ctx.load_cert_chain(os.path.join(CERT_PATH, "server.pem")) + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(2) + self.addCleanup(listener.close) + + 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. + client_ctx = get_ssl_context(None, None, CA_PEM, None, False, False, False, _IS_SYNC) + options = PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=client_ctx) + + created = [] + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + created.append(sock) + return sock + + async def callback(context): + if _IS_SYNC: + return connect() + return await asyncio.get_running_loop().run_in_executor(None, 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 = socket.socketpair() diff --git a/test/test_encryption.py b/test/test_encryption.py index 01cf8cfc83..2f9ed382ca 100644 --- a/test/test_encryption.py +++ b/test/test_encryption.py @@ -226,7 +226,7 @@ def test_init_kms_tls_options(self): self.assertEqual(ctx.verify_mode, ssl.CERT_REQUIRED) -# KMS connect callback unit and prose tests live in test_kms_connect.py. +# KMS connect callback tests live in test_kms_connect.py. class TestClientOptions(PyMongoTestCase): @@ -1999,7 +1999,7 @@ def test_invalid_hostname_in_kms_certificate(self): self.client_encrypted.create_data_key("aws", master_key=key) -# KMS connect callback unit and prose tests live in test_kms_connect.py. +# 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 diff --git a/test/test_kms_connect.py b/test/test_kms_connect.py index 25c96191ab..4bbf6570d0 100644 --- a/test/test_kms_connect.py +++ b/test/test_kms_connect.py @@ -4,6 +4,7 @@ import asyncio import dataclasses +import os import socket import ssl import threading @@ -19,7 +20,7 @@ AutoEncryptionOpts, KMSConnectContext, ) -from pymongo.errors import ConfigurationError, EncryptionError, NetworkTimeout +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 ( @@ -29,7 +30,7 @@ _wrap_encryption_errors, ) from test import PyMongoTestCase -from test.helpers_shared import CLIENT_PEM +from test.helpers_shared import CA_PEM, CERT_PATH, CLIENT_PEM from test.test_encryption import OPTS _IS_SYNC = True @@ -144,6 +145,56 @@ def callback(context): 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 = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_ctx.load_cert_chain(os.path.join(CERT_PATH, "server.pem")) + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(2) + self.addCleanup(listener.close) + + 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. + client_ctx = get_ssl_context(None, None, CA_PEM, None, False, False, False, _IS_SYNC) + options = PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=client_ctx) + + created = [] + + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) + created.append(sock) + return sock + + def callback(context): + if _IS_SYNC: + return connect() + return asyncio.get_running_loop().run_in_executor(None, 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 = socket.socketpair() From 5eee12d5e3afa47cf4a47610c65554085b272498 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 09:20:32 -0500 Subject: [PATCH 3/9] PYTHON-5805 Factor shared helpers into test_kms_connect.py Extract the repeated listener/socketpair setup, TLS context construction, and the flavor-aware run-blocking wrapper into shared helpers, and collapse the fixed-value callbacks into a factory. --- test/asynchronous/test_kms_connect.py | 163 +++++++++++++------------- test/test_kms_connect.py | 159 +++++++++++++------------ 2 files changed, 166 insertions(+), 156 deletions(-) diff --git a/test/asynchronous/test_kms_connect.py b/test/asynchronous/test_kms_connect.py index 99258cadda..a96b0368d2 100644 --- a/test/asynchronous/test_kms_connect.py +++ b/test/asynchronous/test_kms_connect.py @@ -37,13 +37,57 @@ 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 flavor).""" + 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(): - return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=None) + 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") @@ -67,41 +111,36 @@ async def callback(context): context.host = "evil.example.com" # type: ignore[misc] 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) + 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 = socket.socketpair() - self.addCleanup(right.close) + 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) - async def callback(context): - return wrapped - with self.assertRaisesRegex(ConfigurationError, "unwrapped"): - await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + 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 = socket.socketpair() - self.addCleanup(left.close) - self.addCleanup(right.close) + 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.example.com", 443), self._pool_options(), callback, 12.5) + conn = await _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 12.5) self.assertIs(conn, left) self.assertEqual(len(received), 1) @@ -111,12 +150,8 @@ async def callback(context): 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) + server_ctx = _tls_server_context() + listener = self._listen() def serve(): try: @@ -127,9 +162,7 @@ def serve(): 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) + options = self._pool_options(_client_tls_context()) def connect(): sock = socket.create_connection(listener.getsockname(), timeout=10) @@ -137,9 +170,7 @@ def connect(): return sock async def callback(context): - if _IS_SYNC: - return connect() - return await asyncio.get_running_loop().run_in_executor(None, connect) + return await _run_blocking(connect) conn = await _connect_kms(listener.getsockname(), options, callback, 10.0) self.addCleanup(conn.close) @@ -150,12 +181,8 @@ async def test_tls_verification_targets_the_kms_host(self): # 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 = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - server_ctx.load_cert_chain(os.path.join(CERT_PATH, "server.pem")) - listener = socket.socket() - listener.bind(("127.0.0.1", 0)) - listener.listen(2) - self.addCleanup(listener.close) + server_ctx = _tls_server_context(os.path.join(CERT_PATH, "server.pem")) + listener = self._listen(2) def serve(): for _ in range(2): @@ -169,8 +196,7 @@ def serve(): threading.Thread(target=serve, daemon=True).start() # Full verification: trusted CA, invalid certs and hostnames rejected. - client_ctx = get_ssl_context(None, None, CA_PEM, None, False, False, False, _IS_SYNC) - options = PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=client_ctx) + options = self._pool_options(_client_tls_context(verify=True)) created = [] @@ -180,9 +206,7 @@ def connect(): return sock async def callback(context): - if _IS_SYNC: - return connect() - return await asyncio.get_running_loop().run_in_executor(None, connect) + return await _run_blocking(connect) port = listener.getsockname()[1] # The cert covers localhost: verifying against the KMS address succeeds. @@ -197,15 +221,12 @@ async def callback(context): 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) + left, _right = self._socketpair() with self.assertRaisesRegex(ConfigurationError, "TransportSocket"): - await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + 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 @@ -214,13 +235,12 @@ async def test_cancelled_tls_wrap_closes_late_socket(self): raise unittest.SkipTest("the cancel-safe wrap is an async path") from pymongo.pool_shared import _close_late_socket - left, right = socket.socketpair() + 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) - self.addCleanup(right.close) async def test_non_coroutine_callback_is_rejected(self): # A plain def must be rejected before it blocks the event loop. @@ -234,7 +254,7 @@ def callback(context): return None with self.assertRaisesRegex(ConfigurationError, "coroutine function"): - await _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + 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): @@ -242,11 +262,8 @@ async def test_unconnected_socket_from_callback_is_rejected(self): 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) + 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. @@ -257,11 +274,8 @@ async def test_datagram_socket_from_callback_is_rejected(self): 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) + 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 @@ -305,7 +319,7 @@ async def callback(context): # 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) + 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 @@ -313,9 +327,7 @@ async def test_csot_deadline_stops_a_hung_callback(self): if _IS_SYNC: raise unittest.SkipTest("the sync API cannot interrupt a callback") - left, right = socket.socketpair() - self.addCleanup(left.close) - self.addCleanup(right.close) + left, _right = self._socketpair() async def hung_callback(context): await asyncio.sleep(0.5) @@ -323,9 +335,7 @@ async def hung_callback(context): with self.assertRaises(NetworkTimeout): with pymongo.timeout(0.1): - await _connect_kms( - ("kms.example.com", 443), self._pool_options(), hung_callback, 10.0 - ) + 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) @@ -337,12 +347,8 @@ async def test_cancelling_kms_connect_closes_the_callback_socket(self): if _IS_SYNC: raise unittest.SkipTest("cancellation is an async-only behavior") - 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) + server_ctx = _tls_server_context() + listener = self._listen() gate = threading.Event() eof = threading.Event() @@ -382,22 +388,21 @@ def stub_server(): threading.Thread(target=stub_server, daemon=True).start() - 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) + options = self._pool_options(_client_tls_context()) socks = [] - async def callback(context): - sock = await asyncio.get_running_loop().run_in_executor( - None, - lambda: socket.create_connection(listener.getsockname(), timeout=10), - ) + 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 flavor returns a socket instead of a coroutine, so both # error codes are needed depending on the flavor being checked. - connect = _connect_kms(listener.getsockname(), options, callback, 10.0) - task = asyncio.ensure_future(connect) # type: ignore[type-var,arg-type] + 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 diff --git a/test/test_kms_connect.py b/test/test_kms_connect.py index 4bbf6570d0..28d9bc6e1e 100644 --- a/test/test_kms_connect.py +++ b/test/test_kms_connect.py @@ -37,13 +37,57 @@ 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 flavor).""" + 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(): - return PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=None) + 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") @@ -67,41 +111,34 @@ def callback(context): context.host = "evil.example.com" # type: ignore[misc] 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) + _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 = socket.socketpair() - self.addCleanup(right.close) + 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) - def callback(context): - return wrapped - with self.assertRaisesRegex(ConfigurationError, "unwrapped"): - _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + _connect_kms(_KMS_ADDRESS, self._pool_options(), _callback_returning(wrapped), 10.0) def test_context_receives_host_port_and_timeout(self): received = [] - left, right = socket.socketpair() - self.addCleanup(left.close) - self.addCleanup(right.close) + 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.example.com", 443), self._pool_options(), callback, 12.5) + conn = _connect_kms(_KMS_ADDRESS, self._pool_options(), callback, 12.5) self.assertIs(conn, left) self.assertEqual(len(received), 1) @@ -111,12 +148,8 @@ def callback(context): 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) + server_ctx = _tls_server_context() + listener = self._listen() def serve(): try: @@ -127,9 +160,7 @@ def serve(): 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) + options = self._pool_options(_client_tls_context()) def connect(): sock = socket.create_connection(listener.getsockname(), timeout=10) @@ -137,9 +168,7 @@ def connect(): return sock def callback(context): - if _IS_SYNC: - return connect() - return asyncio.get_running_loop().run_in_executor(None, connect) + return _run_blocking(connect) conn = _connect_kms(listener.getsockname(), options, callback, 10.0) self.addCleanup(conn.close) @@ -150,12 +179,8 @@ def test_tls_verification_targets_the_kms_host(self): # 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 = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) - server_ctx.load_cert_chain(os.path.join(CERT_PATH, "server.pem")) - listener = socket.socket() - listener.bind(("127.0.0.1", 0)) - listener.listen(2) - self.addCleanup(listener.close) + server_ctx = _tls_server_context(os.path.join(CERT_PATH, "server.pem")) + listener = self._listen(2) def serve(): for _ in range(2): @@ -169,8 +194,7 @@ def serve(): threading.Thread(target=serve, daemon=True).start() # Full verification: trusted CA, invalid certs and hostnames rejected. - client_ctx = get_ssl_context(None, None, CA_PEM, None, False, False, False, _IS_SYNC) - options = PoolOptions(connect_timeout=10, socket_timeout=10, ssl_context=client_ctx) + options = self._pool_options(_client_tls_context(verify=True)) created = [] @@ -180,9 +204,7 @@ def connect(): return sock def callback(context): - if _IS_SYNC: - return connect() - return asyncio.get_running_loop().run_in_executor(None, connect) + return _run_blocking(connect) port = listener.getsockname()[1] # The cert covers localhost: verifying against the KMS address succeeds. @@ -197,15 +219,12 @@ def callback(context): 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) + left, _right = self._socketpair() with self.assertRaisesRegex(ConfigurationError, "TransportSocket"): - _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + _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 @@ -214,13 +233,12 @@ def test_cancelled_tls_wrap_closes_late_socket(self): raise unittest.SkipTest("the cancel-safe wrap is an async path") from pymongo.pool_shared import _close_late_socket - left, right = socket.socketpair() + 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) - self.addCleanup(right.close) def test_non_coroutine_callback_is_rejected(self): # A plain def must be rejected before it blocks the event loop. @@ -234,7 +252,7 @@ def callback(context): return None with self.assertRaisesRegex(ConfigurationError, "coroutine function"): - _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + _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): @@ -242,11 +260,8 @@ def test_unconnected_socket_from_callback_is_rejected(self): 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) + _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. @@ -257,11 +272,8 @@ def test_datagram_socket_from_callback_is_rejected(self): 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) + _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 @@ -305,7 +317,7 @@ def callback(context): # Not a ConfigurationError, so kms_request retries it. with self.assertRaises(OSError): - _connect_kms(("kms.example.com", 443), self._pool_options(), callback, 10.0) + _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 @@ -313,9 +325,7 @@ def test_csot_deadline_stops_a_hung_callback(self): if _IS_SYNC: raise unittest.SkipTest("the sync API cannot interrupt a callback") - left, right = socket.socketpair() - self.addCleanup(left.close) - self.addCleanup(right.close) + left, _right = self._socketpair() def hung_callback(context): time.sleep(0.5) @@ -323,7 +333,7 @@ def hung_callback(context): with self.assertRaises(NetworkTimeout): with pymongo.timeout(0.1): - _connect_kms(("kms.example.com", 443), self._pool_options(), hung_callback, 10.0) + _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) @@ -335,12 +345,8 @@ def test_cancelling_kms_connect_closes_the_callback_socket(self): if _IS_SYNC: raise unittest.SkipTest("cancellation is an async-only behavior") - 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) + server_ctx = _tls_server_context() + listener = self._listen() gate = threading.Event() eof = threading.Event() @@ -380,22 +386,21 @@ def stub_server(): threading.Thread(target=stub_server, daemon=True).start() - 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) + options = self._pool_options(_client_tls_context()) socks = [] - def callback(context): - sock = asyncio.get_running_loop().run_in_executor( - None, - lambda: socket.create_connection(listener.getsockname(), timeout=10), - ) + def connect(): + sock = socket.create_connection(listener.getsockname(), timeout=10) socks.append(sock) return sock + def callback(context): + return _run_blocking(connect) + # The sync flavor returns a socket instead of a coroutine, so both # error codes are needed depending on the flavor being checked. - connect = _connect_kms(listener.getsockname(), options, callback, 10.0) - task = asyncio.ensure_future(connect) # type: ignore[type-var,arg-type] + 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 From 0870b51be60e126e29e09611000551110b2c49d7 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 09:24:13 -0500 Subject: [PATCH 4/9] PYTHON-5805 Clarify _close_rejected_kms_socket docstring --- pymongo/asynchronous/encryption.py | 7 ++++--- pymongo/synchronous/encryption.py | 7 ++++--- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 0eea43d898..80cc292ffa 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -128,10 +128,11 @@ def _close_rejected_kms_socket(obj: Any) -> None: - """Close a rejected kms_connect_callback return value, best effort. + """Close a callback return value that failed validation, best effort. - Nothing else will close it: _connect_kms raises before the result reaches - the caller's ``finally``. + ``_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): diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index b89ad933f1..0b2e68f663 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -128,10 +128,11 @@ def _close_rejected_kms_socket(obj: Any) -> None: - """Close a rejected kms_connect_callback return value, best effort. + """Close a callback return value that failed validation, best effort. - Nothing else will close it: _connect_kms raises before the result reaches - the caller's ``finally``. + ``_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): From eadebcaeb05d9604a7ad923bdd1ba76ab8ad8462 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 09:28:13 -0500 Subject: [PATCH 5/9] PYTHON-5805 Simplify the result: Any comment --- pymongo/asynchronous/encryption.py | 3 +-- pymongo/synchronous/encryption.py | 3 +-- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 80cc292ffa..a8891fc576 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -169,8 +169,7 @@ async def _connect_kms( raise ConfigurationError( "kms_connect_callback must be a coroutine function for the async API." ) - # Typed as Any so the generated synchronous flavor type-checks: the sync - # callback returns a plain socket, which is not awaitable. + # Any: the generated sync flavor's callback returns a plain socket. result: Any = kms_connect_callback( KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) ) diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index 0b2e68f663..f2b114d3e6 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -169,8 +169,7 @@ def _connect_kms( raise ConfigurationError( "kms_connect_callback must be a coroutine function for the async API." ) - # Typed as Any so the generated synchronous flavor type-checks: the sync - # callback returns a plain socket, which is not awaitable. + # Any: the generated sync flavor's callback returns a plain socket. result: Any = kms_connect_callback( KMSConnectContext(host=address[0], port=cast(int, address[1]), timeout=timeout) ) From 07b62ad85b2a7a5372898c938f7918963b550dd0 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 09:30:25 -0500 Subject: [PATCH 6/9] PYTHON-5805 Tighten the kms_connect_callback docstring --- pymongo/asynchronous/encryption.py | 17 ++++++++--------- pymongo/synchronous/encryption.py | 17 ++++++++--------- tools/synchro.py | 11 +++-------- 3 files changed, 19 insertions(+), 26 deletions(-) diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index a8891fc576..4aab64978e 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -709,15 +709,14 @@ def __init__( 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`, over - which the driver performs the KMS TLS handshake. The callback - must be a coroutine function for the asynchronous API; a plain callable is rejected before it can block the event loop. - When a CSOT timeout is active, the driver stops waiting at the - deadline and closes any socket the callback yields later. - Defaults to ``None``, meaning the driver connects to KMS hosts - directly. + 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. + On timeout the driver stops waiting and closes any late-yielded + socket. Defaults to ``None``, meaning the driver connects to KMS + hosts directly. .. versionchanged:: 4.19 Added the `kms_connect_callback` parameter. diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index f2b114d3e6..aa4641303d 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -707,15 +707,14 @@ def __init__( 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`, over - which the driver performs the KMS TLS handshake. The callback - must be a regular function; the async API requires a coroutine function and rejects plain callables before they can block the event loop. - When a CSOT timeout is active, the driver stops waiting at the - deadline and closes any socket the callback yields later. - Defaults to ``None``, meaning the driver connects to KMS hosts - directly. + 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. + On timeout the driver stops waiting and closes any late-yielded + socket. Defaults to ``None``, meaning the driver connects to KMS + hosts directly. .. versionchanged:: 4.19 Added the `kms_connect_callback` parameter. diff --git a/tools/synchro.py b/tools/synchro.py index bd5edebda3..b5c6b2873a 100644 --- a/tools/synchro.py +++ b/tools/synchro.py @@ -301,15 +301,10 @@ def translate_docstrings(lines: list[str]) -> list[str]: 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; a plain callable is rejected before it can block the event loop" - in lines[i] - ): + 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; a plain callable is rejected before it can block the event loop", - "must be a regular function; the async API requires a " - "coroutine function and rejects plain callables before " - "they can block the event loop", + "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]: From 089ffafa86267e0e3c3ca1c7e0551d16c1e09215 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 09:33:23 -0500 Subject: [PATCH 7/9] PYTHON-5805 Fix sync docstring's false deadline-enforcement promise The synchronous API waits for the callback without enforcing the CSOT deadline; state the flavor-neutral contract instead (callback honors context.timeout). Addresses Copilot review discussion_r4185118746. --- pymongo/asynchronous/encryption.py | 6 +++--- pymongo/synchronous/encryption.py | 6 +++--- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 4aab64978e..30a0e6d574 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -714,9 +714,9 @@ def __init__( 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. - On timeout the driver stops waiting and closes any late-yielded - socket. Defaults to ``None``, meaning the driver connects to KMS - hosts directly. + 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. diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index aa4641303d..e22cc50405 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -712,9 +712,9 @@ def __init__( returns a connected, unwrapped :class:`socket.socket`, over which the driver performs the KMS TLS handshake. Must be a regular function. - On timeout the driver stops waiting and closes any late-yielded - socket. Defaults to ``None``, meaning the driver connects to KMS - hosts directly. + 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. From fa6adfec4c8703f1e533fc4faf48ca319d9b90e4 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 10:28:11 -0500 Subject: [PATCH 8/9] PYTHON-5805 Reword the result: Any comment --- pymongo/asynchronous/encryption.py | 2 +- pymongo/synchronous/encryption.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/pymongo/asynchronous/encryption.py b/pymongo/asynchronous/encryption.py index 30a0e6d574..4ae63716e3 100644 --- a/pymongo/asynchronous/encryption.py +++ b/pymongo/asynchronous/encryption.py @@ -169,7 +169,7 @@ async def _connect_kms( raise ConfigurationError( "kms_connect_callback must be a coroutine function for the async API." ) - # Any: the generated sync flavor's callback returns a plain socket. + # 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) ) diff --git a/pymongo/synchronous/encryption.py b/pymongo/synchronous/encryption.py index e22cc50405..ac2f710382 100644 --- a/pymongo/synchronous/encryption.py +++ b/pymongo/synchronous/encryption.py @@ -169,7 +169,7 @@ def _connect_kms( raise ConfigurationError( "kms_connect_callback must be a coroutine function for the async API." ) - # Any: the generated sync flavor's callback returns a plain socket. + # 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) ) From d76bf9938623f30a98ba323fcceb3842055a6046 Mon Sep 17 00:00:00 2001 From: Steven Silvester Date: Mon, 5 Oct 2026 10:30:09 -0500 Subject: [PATCH 9/9] PYTHON-5805 Say sync version instead of sync flavor --- test/asynchronous/test_kms_connect.py | 6 +++--- test/test_kms_connect.py | 6 +++--- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/test/asynchronous/test_kms_connect.py b/test/asynchronous/test_kms_connect.py index a96b0368d2..6a577b90ee 100644 --- a/test/asynchronous/test_kms_connect.py +++ b/test/asynchronous/test_kms_connect.py @@ -54,7 +54,7 @@ def _client_tls_context(verify=False): async def _run_blocking(func, *args): - """Run a blocking callable off the event loop (inline in the sync flavor).""" + """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) @@ -399,8 +399,8 @@ def connect(): async def callback(context): return await _run_blocking(connect) - # The sync flavor returns a socket instead of a coroutine, so both - # error codes are needed depending on the flavor being checked. + # 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): diff --git a/test/test_kms_connect.py b/test/test_kms_connect.py index 28d9bc6e1e..8bd59ca689 100644 --- a/test/test_kms_connect.py +++ b/test/test_kms_connect.py @@ -54,7 +54,7 @@ def _client_tls_context(verify=False): def _run_blocking(func, *args): - """Run a blocking callable off the event loop (inline in the sync flavor).""" + """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) @@ -397,8 +397,8 @@ def connect(): def callback(context): return _run_blocking(connect) - # The sync flavor returns a socket instead of a coroutine, so both - # error codes are needed depending on the flavor being checked. + # 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):