diff --git a/doc/changelog.rst b/doc/changelog.rst index 41e6dee233..b937ed0d20 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -10,8 +10,23 @@ Bug fixes - Fixed a bug where the synchronous client could permanently deadlock under gevent when a greenlet was killed while checking a connection back into the pool (`PYTHON-6074`_). +- ``MongoClient.append_metadata()`` and ``AsyncMongoClient.append_metadata()`` + now detect duplicates by comparing the whole + :class:`~pymongo.driver_info.DriverInfo` instead of only its name. The + comparison is exact, so drivers that differ in name case or platform are no + longer treated as duplicates (`PYTHON-6040`_). +- ``driver.name`` and ``driver.version`` in the handshake metadata are now + ``|``-delimited lists with 1:1 index correspondence, including empty version + entries for the built-in ``|c`` and ``|async`` name segments + (`PYTHON-6040`_). +- Fixed a bug where truncating the handshake metadata to 512 bytes could leave + ``driver.name`` and ``driver.version`` with different numbers of ``|`` + delimiters (`PYTHON-6040`_). +- :class:`~pymongo.driver_info.DriverInfo` now raises :class:`ValueError` when + any field contains the reserved ``|`` delimiter (`PYTHON-6040`_). .. _PYTHON-6074: https://jira.mongodb.org/browse/PYTHON-6074 +.. _PYTHON-6040: https://jira.mongodb.org/browse/PYTHON-6040 Changes in Version 4.18.1 (2026/09/10) -------------------------------------- diff --git a/pymongo/driver_info.py b/pymongo/driver_info.py index 18a51ae638..a325ea1493 100644 --- a/pymongo/driver_info.py +++ b/pymongo/driver_info.py @@ -29,8 +29,14 @@ class DriverInfo(namedtuple("DriverInfo", ["name", "version", "platform"])): The MongoDB server logs PyMongo's name, version, and platform whenever PyMongo establishes a connection. A driver implemented on top of PyMongo can add its own info to this log message. Initialize with three strings - like 'MyDriver', '1.2.3', 'some platform info'. Any of these strings may be - None to accept PyMongo's default. + like 'MyDriver', '1.2.3', 'some platform info'. Any of these strings may + be None. A None or empty name or version appends an empty metadata + segment, keeping the ``driver.name`` and ``driver.version`` segment + counts aligned. A None or empty platform omits the platform. + + The ``|`` character is the reserved delimiter used to join appended + metadata, so it must not appear in any of the fields. A + :class:`ValueError` is raised if it does. """ def __new__( @@ -42,5 +48,7 @@ def __new__( raise TypeError( f"Wrong type for DriverInfo {key} option, value must be an instance of str, not {type(value)}" ) + if value and "|" in value: + raise ValueError(f"DriverInfo {key} must not contain the '|' delimiter") return self diff --git a/pymongo/pool_options.py b/pymongo/pool_options.py index 8b26b4baf2..6d1cc61e9a 100644 --- a/pymongo/pool_options.py +++ b/pymongo/pool_options.py @@ -24,6 +24,7 @@ import platform import sys from collections.abc import MutableMapping +from contextlib import AbstractContextManager, nullcontext from pathlib import Path from typing import TYPE_CHECKING, Any, Optional @@ -37,6 +38,7 @@ WAIT_QUEUE_TIMEOUT, has_c, ) +from pymongo.lock import _create_lock if TYPE_CHECKING: from pymongo.auth_shared import MongoCredential @@ -200,60 +202,117 @@ def _metadata_env() -> dict[str, Any]: _MAX_METADATA_SIZE = 512 +def _truncate_utf8(content: str, overflow: int) -> str: + """Trim `overflow` UTF-8 bytes from the end of content, keeping a valid prefix.""" + if overflow <= 0: + return content + data = content.encode("utf-8") + if len(data) <= overflow: + return "" + return data[: len(data) - overflow].decode("utf-8", errors="ignore") + + +def _normalize_driver(driver: DriverInfo) -> DriverInfo: + """Treat None and "" as equivalent unset fields for deduplication.""" + return driver._replace( + name=driver.name or "", + version=driver.version or "", + platform=driver.platform or "", + ) + + +def _element_size(key: str, value: Any) -> int: + """Size in bytes of the BSON element for ``key``, excluding doc overhead.""" + # bson.encode({key: value}) is 4 bytes of document header plus 1 byte of + # document terminator on top of the element itself. + return len(bson.encode({key: value})) - 5 + + # See: https://github.com/mongodb/specifications/blob/master/source/mongodb-handshake/handshake.md#limitations def _truncate_metadata(metadata: MutableMapping[str, Any]) -> None: """Perform metadata truncation.""" - if len(bson.encode(metadata)) <= _MAX_METADATA_SIZE: + # The only full encode is the initial size check; each step then shrinks + # the tracked size by the exact bytes its change removes. + size = len(bson.encode(metadata)) + if size <= _MAX_METADATA_SIZE: return # 1. Omit fields from env except env.name. env_name = metadata.get("env", {}).get("name") if env_name: - metadata["env"] = {"name": env_name} - if len(bson.encode(metadata)) <= _MAX_METADATA_SIZE: + env = {"name": env_name} + size += _element_size("env", env) - _element_size("env", metadata["env"]) + metadata["env"] = env + if size <= _MAX_METADATA_SIZE: return # 2. Omit fields from os except os.type. os_type = metadata.get("os", {}).get("type") if os_type: - metadata["os"] = {"type": os_type} - if len(bson.encode(metadata)) <= _MAX_METADATA_SIZE: + old_os = metadata["os"] + new_os = {"type": os_type} + size += _element_size("os", new_os) - _element_size("os", old_os) + metadata["os"] = new_os + if size <= _MAX_METADATA_SIZE: return # 3. Omit the env document entirely. - metadata.pop("env", None) - encoded_size = len(bson.encode(metadata)) - if encoded_size <= _MAX_METADATA_SIZE: + env = metadata.pop("env", None) + if env is not None: + size -= _element_size("env", env) + if size <= _MAX_METADATA_SIZE: return # 4. Truncate platform. - overflow = encoded_size - _MAX_METADATA_SIZE + overflow = size - _MAX_METADATA_SIZE plat = metadata.get("platform", "") if plat: - plat = plat[:-overflow] - if plat: - metadata["platform"] = plat + truncated = _truncate_utf8(plat, overflow) + if truncated: + size += _element_size("platform", truncated) - _element_size("platform", plat) + else: + size -= _element_size("platform", plat) + if truncated: + metadata["platform"] = truncated + else: + del metadata["platform"] else: - metadata.pop("platform", None) - encoded_size = len(bson.encode(metadata)) - if encoded_size <= _MAX_METADATA_SIZE: - return - # 5. Truncate driver info. - overflow = encoded_size - _MAX_METADATA_SIZE + # The platform field may be present but empty. + plat = metadata.pop("platform", None) + if plat is not None: + size -= _element_size("platform", plat) + # 5. Truncate driver info, keeping name and version 1:1 index-aligned. driver = metadata.get("driver", {}) if driver: - # Truncate driver version. - driver_version = driver.get("version")[:-overflow] - if len(driver_version) >= len(_METADATA["driver"]["version"]): - metadata["driver"]["version"] = driver_version - else: - metadata["driver"]["version"] = _METADATA["driver"]["version"] - encoded_size = len(bson.encode(metadata)) - if encoded_size <= _MAX_METADATA_SIZE: - return - # Truncate driver name. - overflow = encoded_size - _MAX_METADATA_SIZE - driver_name = driver.get("name")[:-overflow] - if len(driver_name) >= len(_METADATA["driver"]["name"]): - metadata["driver"]["name"] = driver_name - else: - metadata["driver"]["name"] = _METADATA["driver"]["name"] + # Keep the name and version segments paired so they stay 1:1 aligned, + # trimming wrapper content and dropping paired segments only as a + # last resort. Blank pairs (for example from a platform-only appended + # driver) are dropped first since they cost nothing. + name_str = driver.get("name", "") + version_str = driver.get("version", "") + raw_pairs = list(zip(name_str.split("|"), version_str.split("|"))) + pairs = [(name, version) for name, version in raw_pairs if name or version] + new_name = "|".join(name for name, _ in pairs) + new_version = "|".join(version for _, version in pairs) + size += (len(new_name.encode("utf-8")) - len(name_str.encode("utf-8"))) + ( + len(new_version.encode("utf-8")) - len(version_str.encode("utf-8")) + ) + driver["name"] = new_name + driver["version"] = new_version + overflow = size - _MAX_METADATA_SIZE + while overflow > 0 and len(pairs) > 1: + # A single remaining pair never exceeds the limit: it is the base + # driver pair, which steps 1-4 left small enough to fit. + last_name, last_version = pairs[-1] + if last_version: + new_version = _truncate_utf8(last_version, overflow) + overflow -= len(last_version.encode("utf-8")) - len(new_version.encode("utf-8")) + pairs[-1] = (last_name, new_version) + elif last_name: + new_name = _truncate_utf8(last_name, overflow) + overflow -= len(last_name.encode("utf-8")) - len(new_name.encode("utf-8")) + pairs[-1] = (new_name, last_version) + else: + pairs.pop() + overflow -= len(f"|{last_name}|{last_version}".encode()) + driver["name"] = "|".join(name for name, _ in pairs) + driver["version"] = "|".join(version for _, version in pairs) # If the first getaddrinfo call of this interpreter's life is on a thread, @@ -277,6 +336,7 @@ class PoolOptions: """ __slots__ = ( + "__appended_drivers", "__appname", "__compression_settings", "__connect_timeout", @@ -288,6 +348,7 @@ class PoolOptions: "__max_idle_time_seconds", "__max_pool_size", "__metadata", + "__metadata_lock", "__min_pool_size", "__pause_enabled", "__server_api", @@ -336,6 +397,11 @@ def __init__( self.__load_balanced = load_balanced self.__credentials = credentials self.__metadata = copy.deepcopy(_METADATA) + self.__appended_drivers: set[DriverInfo] = set() + # Only the synchronous client can append metadata from multiple threads. + self.__metadata_lock: AbstractContextManager[bool | None] = ( + _create_lock() if is_sync else nullcontext() + ) if appname: self.__metadata["application"] = {"name": appname} @@ -353,11 +419,19 @@ def __init__( self.__metadata["driver"]["name"], "c", ) + self.__metadata["driver"]["version"] = "{}|{}".format( + self.__metadata["driver"]["version"], + "", + ) if not is_sync: self.__metadata["driver"]["name"] = "{}|{}".format( self.__metadata["driver"]["name"], "async", ) + self.__metadata["driver"]["version"] = "{}|{}".format( + self.__metadata["driver"]["version"], + "", + ) if driver: self._update_metadata(driver) @@ -368,28 +442,40 @@ def __init__( _truncate_metadata(self.__metadata) def _update_metadata(self, driver: DriverInfo) -> None: - """Updates the client's metadata""" - if driver.name and driver.name.lower() in self.__metadata["driver"]["name"].lower().split( - "|" - ): - return - - metadata = copy.deepcopy(self.__metadata) - - if driver.name: - metadata["driver"]["name"] = "{}|{}".format( - metadata["driver"]["name"], - driver.name, - ) - if driver.version: + """Updates the client's metadata.""" + with self.__metadata_lock: + driver = _normalize_driver(driver) + if driver in self.__appended_drivers: + return + + name_delims = self.__metadata["driver"]["name"].count("|") + version_delims = self.__metadata["driver"]["version"].count("|") + # Only the top-level keys and the "driver" document are mutated, + # so shallow copies of those two are enough. + metadata = {**self.__metadata, "driver": dict(self.__metadata["driver"])} + + metadata["driver"]["name"] = "{}|{}".format(metadata["driver"]["name"], driver.name) metadata["driver"]["version"] = "{}|{}".format( - metadata["driver"]["version"], - driver.version, + metadata["driver"]["version"], driver.version ) - if driver.platform: - metadata["platform"] = "{}|{}".format(metadata["platform"], driver.platform) - - self.__metadata = metadata + if driver.platform: + if "platform" in metadata: + metadata["platform"] = "{}|{}".format(metadata["platform"], driver.platform) + else: + metadata["platform"] = driver.platform + + _truncate_metadata(metadata) + + self.__metadata = metadata + + # Only track drivers whose appended name/version pair survived + # truncation (i.e. both gained a segment), so __appended_drivers + # stays bounded. + if ( + metadata["driver"]["name"].count("|") > name_delims + and metadata["driver"]["version"].count("|") > version_delims + ): + self.__appended_drivers.add(driver) @property def _credentials(self) -> Optional[MongoCredential]: diff --git a/test/asynchronous/test_client.py b/test/asynchronous/test_client.py index 90a2d33a45..fcad59a563 100644 --- a/test/asynchronous/test_client.py +++ b/test/asynchronous/test_client.py @@ -90,7 +90,13 @@ WriteConcernError, ) from pymongo.monitoring import ServerHeartbeatListener, ServerHeartbeatStartedEvent -from pymongo.pool_options import _MAX_METADATA_SIZE, _METADATA, ENV_VAR_K8S, PoolOptions +from pymongo.pool_options import ( + _MAX_METADATA_SIZE, + _METADATA, + ENV_VAR_K8S, + PoolOptions, + _truncate_metadata, +) from pymongo.read_preferences import ReadPreference from pymongo.server_description import ServerDescription from pymongo.server_selectors import readable_server_selector, writable_server_selector @@ -124,6 +130,8 @@ NTHREADS, CMAPListener, FunctionCallRecorder, + _driver_version, + _metadata_with_appended_driver, delay, gevent_monkey_patched, is_greenthread_patched, @@ -386,6 +394,9 @@ async def test_metadata(self): metadata["driver"]["name"] = "PyMongo|c|async" else: metadata["driver"]["name"] = "PyMongo|async" + metadata["driver"]["version"] = _driver_version( + _METADATA["driver"]["version"], metadata["driver"]["name"] + ) metadata["application"] = {"name": "foobar"} client = self.simple_client("mongodb://foo:27017/?appname=foobar&connect=false") options = client.options @@ -397,6 +408,8 @@ async def test_metadata(self): self.simple_client(appname="x" * 128) with self.assertRaises(ValueError): self.simple_client(appname="x" * 129) + + async def test_metadata_bad_driver_options(self): # Bad "driver" options. self.assertRaises(TypeError, DriverInfo, "Foo", 1, "a") self.assertRaises(TypeError, DriverInfo, version="1", platform="a") @@ -407,12 +420,9 @@ async def test_metadata(self): self.simple_client(driver="abc") with self.assertRaises(TypeError): self.simple_client(driver=("Foo", "1", "a")) - # Test appending to driver info. - if has_c(): - metadata["driver"]["name"] = "PyMongo|c|async|FooDriver" - else: - metadata["driver"]["name"] = "PyMongo|async|FooDriver" - metadata["driver"]["version"] = "{}|1.2.3".format(_METADATA["driver"]["version"]) + + async def test_metadata_appends_driver_info(self): + metadata = _metadata_with_appended_driver(_IS_SYNC, "FooDriver", "1.2.3") client = self.simple_client( "foo", 27017, @@ -420,9 +430,9 @@ async def test_metadata(self): driver=DriverInfo("FooDriver", "1.2.3", None), connect=False, ) - options = client.options - self.assertEqual(options.pool_options.metadata, metadata) - metadata["platform"] = "{}|FooPlatform".format(_METADATA["platform"]) + self.assertEqual(client.options.pool_options.metadata, metadata) + + metadata = _metadata_with_appended_driver(_IS_SYNC, "FooDriver", "1.2.3", "FooPlatform") client = self.simple_client( "foo", 27017, @@ -430,27 +440,131 @@ async def test_metadata(self): driver=DriverInfo("FooDriver", "1.2.3", "FooPlatform"), connect=False, ) - options = client.options - self.assertEqual(options.pool_options.metadata, metadata) - # Test truncating driver info metadata. + self.assertEqual(client.options.pool_options.metadata, metadata) + + async def test_metadata_truncates_driver_info(self): + # Truncated driver info must stay within the limit and keep name and + # version index-aligned. client = self.simple_client( driver=DriverInfo(name="s" * _MAX_METADATA_SIZE), connect=False, ) options = client.options + truncated = options.pool_options.metadata["driver"] self.assertLessEqual( len(bson.encode(options.pool_options.metadata)), _MAX_METADATA_SIZE, ) + self.assertEqual( + truncated["name"].count("|"), + truncated["version"].count("|"), + ) client = self.simple_client( driver=DriverInfo(name="s" * _MAX_METADATA_SIZE, version="s" * _MAX_METADATA_SIZE), connect=False, ) options = client.options + truncated = options.pool_options.metadata["driver"] self.assertLessEqual( len(bson.encode(options.pool_options.metadata)), _MAX_METADATA_SIZE, ) + self.assertEqual( + truncated["name"].count("|"), + truncated["version"].count("|"), + ) + # An oversized wrapper name with no version must retain a truncated + # name rather than collapse to the base entry. + client = self.simple_client( + driver=DriverInfo(name="x" * (_MAX_METADATA_SIZE * 2), version=None), + connect=False, + ) + truncated = client.options.pool_options.metadata["driver"] + self.assertLessEqual( + len(bson.encode(client.options.pool_options.metadata)), + _MAX_METADATA_SIZE, + ) + self.assertIn("xxxx", truncated["name"]) + self.assertEqual( + truncated["name"].count("|"), + truncated["version"].count("|"), + ) + + async def test_metadata_truncation_drops_blank_pairs_first(self): + # A platform-only driver appends a blank name/version pair. Truncation + # must drop that pair before trimming a real driver's content. + def metadata_for(name_len: int) -> dict[str, Any]: + name = "PyMongo||" + "W" * name_len # blank pair in the middle + return {"driver": {"name": name, "version": "1.0||1.0"}} + + # Size the name so the document is exactly two bytes over the limit: + # dropping the blank pair alone must make it fit. + name_len = next( + n + for n in range(1, 2 * _MAX_METADATA_SIZE) + if len(bson.encode(metadata_for(n))) == _MAX_METADATA_SIZE + 2 + ) + metadata = metadata_for(name_len) + _truncate_metadata(metadata) + self.assertEqual(len(bson.encode(metadata)), _MAX_METADATA_SIZE) + self.assertEqual( + metadata["driver"], + {"name": "PyMongo|" + "W" * name_len, "version": "1.0|1.0"}, + ) + + async def test_metadata_append_is_bounded(self): + # Successive appends must stay within the limit and keep name and + # version index-aligned after truncation. Once the metadata saturates, + # further appends must not grow the dedup tracking list. + client = self.simple_client(connect=False) + for i in range(300): + client.append_metadata(DriverInfo(name=f"D{i}", version=f"1.{i}")) + pool = client.options.pool_options + self.assertLessEqual(len(bson.encode(pool.metadata)), _MAX_METADATA_SIZE) + self.assertEqual( + pool.metadata["driver"]["name"].count("|"), + pool.metadata["driver"]["version"].count("|"), + ) + count = len(pool._PoolOptions__appended_drivers) + for i in range(300, 600): + client.append_metadata(DriverInfo(name=f"D{i}", version=f"1.{i}")) + self.assertEqual(len(pool._PoolOptions__appended_drivers), count) + # Platform-only appends (empty name/version) stay bounded the same way. + client = self.simple_client(connect=False) + for i in range(300): + client.append_metadata(DriverInfo(name="", version="", platform=f"P{i}")) + pool = client.options.pool_options + self.assertLessEqual(len(bson.encode(pool.metadata)), _MAX_METADATA_SIZE) + count = len(pool._PoolOptions__appended_drivers) + for i in range(300, 600): + client.append_metadata(DriverInfo(name="", version="", platform=f"P{i}")) + self.assertEqual(len(pool._PoolOptions__appended_drivers), count) + + async def test_metadata_recreates_platform_after_truncation(self): + # Appending a platform after truncation has dropped it recreates the field. + client = self.simple_client(connect=False) + for i in range(300): + client.append_metadata(DriverInfo(name="", version="", platform=f"Q{i}")) + pool = client.options.pool_options + self.assertLess(len(pool._PoolOptions__appended_drivers), 300) + client.append_metadata(DriverInfo(name="Wrapper", version="1.0", platform="Recreated")) + self.assertLessEqual(len(bson.encode(pool.metadata)), _MAX_METADATA_SIZE) + self.assertEqual( + pool.metadata["driver"]["name"].count("|"), + pool.metadata["driver"]["version"].count("|"), + ) + + async def test_metadata_deduplicates_none_and_empty(self): + # Empty strings are treated as unset, so a duplicate differing only in + # None vs "" is a no-op. + client = self.simple_client(connect=False) + client.append_metadata(DriverInfo("library", None, "Library Platform")) + names = client.options.pool_options.metadata["driver"]["name"] + vers = client.options.pool_options.metadata["driver"]["version"] + client.append_metadata(DriverInfo("library", "", "Library Platform")) + metadata = client.options.pool_options.metadata + self.assertEqual(metadata["driver"]["name"], names) + self.assertEqual(metadata["driver"]["version"], vers) @mock.patch.dict("os.environ", {ENV_VAR_K8S: "1"}) def test_container_metadata(self): @@ -2224,6 +2338,9 @@ async def _test_handshake(self, env_vars, expected_env): metadata["driver"]["name"] = "PyMongo|c|async" else: metadata["driver"]["name"] = "PyMongo|async" + metadata["driver"]["version"] = _driver_version( + _METADATA["driver"]["version"], metadata["driver"]["name"] + ) if expected_env is not None: metadata["env"] = expected_env diff --git a/test/asynchronous/test_client_metadata.py b/test/asynchronous/test_client_metadata.py index 1a07e835a8..56c72d52d6 100644 --- a/test/asynchronous/test_client_metadata.py +++ b/test/asynchronous/test_client_metadata.py @@ -18,7 +18,7 @@ import pathlib import time import unittest -from typing import Any, Optional +from typing import Any, Optional, cast import pytest @@ -99,20 +99,14 @@ async def check_metadata_added( new_name, new_version, new_platform, new_metadata = await self.send_ping_and_get_metadata( client, True ) - if add_name is not None and add_name.lower() in name.lower().split("|"): - self.assertEqual(name, new_name) - self.assertEqual(version, new_version) - self.assertEqual(platform, new_platform) - else: - self.assertEqual(new_name, f"{name}|{add_name}" if add_name is not None else name) - self.assertEqual( - new_version, - f"{version}|{add_version}" if add_version is not None else version, - ) - self.assertEqual( - new_platform, - f"{platform}|{add_platform}" if add_platform is not None else platform, - ) + # Name and version always get a delimiter (empty string if None) to + # preserve 1:1 index correspondence. + self.assertEqual(new_name, f"{name}|{add_name or ''}") + self.assertEqual(new_version, f"{version}|{add_version or ''}") + self.assertEqual( + new_platform, + f"{platform}|{add_platform}" if add_platform is not None else platform, + ) metadata.pop("driver") metadata.pop("platform") @@ -120,7 +114,7 @@ async def check_metadata_added( new_metadata.pop("platform") self.assertEqual(metadata, new_metadata) - async def test_append_metadata(self): + async def test_1_test_that_the_driver_updates_metadata(self): client = await self.async_rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -128,7 +122,7 @@ async def test_append_metadata(self): ) await self.check_metadata_added(client, "framework", "2.0", "Framework Platform") - async def test_append_metadata_platform_none(self): + async def test_1_test_that_the_driver_updates_metadata_platform_none(self): client = await self.async_rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -136,7 +130,7 @@ async def test_append_metadata_platform_none(self): ) await self.check_metadata_added(client, "framework", "2.0", None) - async def test_append_metadata_version_none(self): + async def test_1_test_that_the_driver_updates_metadata_version_none(self): client = await self.async_rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -144,7 +138,7 @@ async def test_append_metadata_version_none(self): ) await self.check_metadata_added(client, "framework", None, "Framework Platform") - async def test_append_metadata_platform_version_none(self): + async def test_1_test_that_the_driver_updates_metadata_platform_version_none(self): client = await self.async_rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -152,14 +146,14 @@ async def test_append_metadata_platform_version_none(self): ) await self.check_metadata_added(client, "framework", None, None) - async def test_multiple_successive_metadata_updates(self): + async def test_2_multiple_successive_metadata_updates(self): client = await self.async_rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, connect=False ) client.append_metadata(DriverInfo("library", "1.2", "Library Platform")) await self.check_metadata_added(client, "framework", "2.0", "Framework Platform") - async def test_multiple_successive_metadata_updates_platform_none(self): + async def test_2_multiple_successive_metadata_updates_platform_none(self): client = await self.async_rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -167,7 +161,7 @@ async def test_multiple_successive_metadata_updates_platform_none(self): client.append_metadata(DriverInfo("library", "1.2", "Library Platform")) await self.check_metadata_added(client, "framework", "2.0", None) - async def test_multiple_successive_metadata_updates_version_none(self): + async def test_2_multiple_successive_metadata_updates_version_none(self): client = await self.async_rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -175,7 +169,7 @@ async def test_multiple_successive_metadata_updates_version_none(self): client.append_metadata(DriverInfo("library", "1.2", "Library Platform")) await self.check_metadata_added(client, "framework", None, "Framework Platform") - async def test_multiple_successive_metadata_updates_platform_version_none(self): + async def test_2_multiple_successive_metadata_updates_platform_version_none(self): client = await self.async_rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -216,10 +210,16 @@ async def test_duplicate_driver_name_no_op(self): await self.check_metadata_added(client, "framework", None, None) # wait for connection to become idle await asyncio.sleep(0.005) - # add same metadata again - await self.check_metadata_added(client, "Framework", None, None) + # Append the exact same DriverInfo again: no-op. + name, version, platform, _ = await self.send_ping_and_get_metadata(client, True) + await asyncio.sleep(0.005) + client.append_metadata(DriverInfo("framework", None, None)) + new_name, new_version, new_platform, _ = await self.send_ping_and_get_metadata(client, True) + self.assertEqual(new_name, name) + self.assertEqual(new_version, version) + self.assertEqual(new_platform, platform) - async def test_handshake_documents_include_backpressure(self): + async def test_9_handshake_documents_include_backpressure(self): # Create a `MongoClient` that is configured to record all handshake documents sent to the server as a part of # connection establishment. client = await self.async_rs_or_single_client("mongodb://" + self.server.address_string) @@ -232,6 +232,125 @@ async def test_handshake_documents_include_backpressure(self): # the document has a field `backpressure` whose value is `"2"`. self.assertEqual(self.handshake_req["backpressure"], "2") + async def test_10_entries_in_driver_name_and_driver_version_correspond_by_index(self): + cases = [ + ("Gap in middle (name)", [(None, None), ("F2", None)], "||F2", "||"), + ("Gap in middle (version)", [("F1", None), ("F2", "2.0")], "|F1|F2", "||2.0"), + ("Trailing delimiter retained", [("F1", None)], "|F1", "|"), + ( + "Equal versions do not collapse", + [("F1", "{driver_version}")], + "|F1", + "|{driver_version}", + ), + ( + "Equal names do not collapse", + [("{driver_name}", "1.0")], + "|{driver_name}", + "|1.0", + ), + ("Duplicates still deduplicate", [("F1", "1.0"), ("F1", "1.0")], "|F1", "|1.0"), + ("All versions absent", [("F1", None), ("F2", None)], "|F1|F2", "||"), + ("All names absent", [(None, "1.0"), (None, "2.0")], "||", "|1.0|2.0"), + ( + "Non-adjacent duplicate", + [("F1", "1.0"), ("F2", "2.0"), ("F1", "1.0")], + "|F1|F2", + "|1.0|2.0", + ), + ( + "Platform-only difference is not a duplicate", + [("F1", "1.0", "P1"), ("F1", "1.0", "P2")], + "|F1|F1", + "|1.0|1.0", + ), + ( + "Wrapper matching the driver's own identity", + [("{driver_name}", "{driver_version}")], + "|{driver_name}", + "|{driver_version}", + ), + ] + for ( + description, + appended, + expected_name_suffix, + expected_version_suffix, + ) in cases: + with self.subTest(description=description): + client = await self.async_rs_or_single_client( + "mongodb://" + self.server.address_string, + maxIdleTimeMS=1, + ) + self.addAsyncCleanup(client.close) + # Capture the driver's own name and version from the first handshake. + name0, version0, _, _ = await self.send_ping_and_get_metadata(client, True) + await asyncio.sleep(0.005) + + self.assertIsNotNone(name0) + self.assertIsNotNone(version0) + version0 = cast(str, version0) + driver_name = name0.split("|")[0] + driver_version = version0.split("|")[0] + + def resolve(value: Optional[str]) -> Optional[str]: + if value is None: + return None + return value.format(driver_name=driver_name, driver_version=driver_version) + + # Append each DriverInfo in order. + for opts in appended: + d_name = resolve(opts[0]) if len(opts) > 0 else None + d_version = resolve(opts[1]) if len(opts) > 1 else None + d_platform = resolve(opts[2]) if len(opts) > 2 else None + client.append_metadata(DriverInfo(d_name or "", d_version, d_platform)) + + # New handshake with the appended metadata. + name1, version1, _, _ = await self.send_ping_and_get_metadata(client, True) + + self.assertEqual( + name1, + name0 + + expected_name_suffix.format( + driver_name=driver_name, driver_version=driver_version + ), + ) + self.assertEqual( + version1, + version0 + + expected_version_suffix.format( + driver_name=driver_name, driver_version=driver_version + ), + ) + + async def test_11_appending_metadata_containing_the_delimiter_raises_an_error(self): + cases = [ + ("frame|work", "2.0", "Framework Platform"), + ("framework", "2|0", "Framework Platform"), + ("framework", "2.0", "Framework|Platform"), + ] + for name, version, platform in cases: + with self.subTest(name=name, version=version, platform=platform): + client = await self.async_rs_or_single_client( + "mongodb://" + self.server.address_string, + maxIdleTimeMS=1, + driver=DriverInfo("library", "1.2", "Library Platform"), + ) + self.addAsyncCleanup(client.close) + # Send initial handshake. + name0, version0, platform0, _metadata = await self.send_ping_and_get_metadata( + client, True + ) + await asyncio.sleep(0.005) + # Constructing metadata containing the delimiter raises. + with self.assertRaises(ValueError): + DriverInfo(name, version, platform) + # Metadata is unchanged on the next handshake. + name1, version1, platform1, _ = await self.send_ping_and_get_metadata(client, True) + self.assertEqual(name1, name0) + self.assertEqual(version1, version0) + self.assertEqual(platform1, platform0) + if __name__ == "__main__": unittest.main() diff --git a/test/mockupdb/test_handshake.py b/test/mockupdb/test_handshake.py index 2772e6f77a..e3a3fc0562 100644 --- a/test/mockupdb/test_handshake.py +++ b/test/mockupdb/test_handshake.py @@ -49,9 +49,11 @@ def _check_handshake_data(request): assert data["application"] == {"name": "my app"} if has_c(): name = "PyMongo|c" + version = pymongo_version + "|" else: name = "PyMongo" - assert data["driver"] == {"name": name, "version": pymongo_version} + version = pymongo_version + assert data["driver"] == {"name": name, "version": version} # Keep it simple, just check these fields exist. assert "os" in data diff --git a/test/test_client.py b/test/test_client.py index 249f95d8fc..91d0c71d13 100644 --- a/test/test_client.py +++ b/test/test_client.py @@ -81,7 +81,13 @@ WriteConcernError, ) from pymongo.monitoring import ServerHeartbeatListener, ServerHeartbeatStartedEvent -from pymongo.pool_options import _MAX_METADATA_SIZE, _METADATA, ENV_VAR_K8S, PoolOptions +from pymongo.pool_options import ( + _MAX_METADATA_SIZE, + _METADATA, + ENV_VAR_K8S, + PoolOptions, + _truncate_metadata, +) from pymongo.read_preferences import ReadPreference from pymongo.server_description import ServerDescription from pymongo.server_selectors import readable_server_selector, writable_server_selector @@ -123,6 +129,8 @@ NTHREADS, CMAPListener, FunctionCallRecorder, + _driver_version, + _metadata_with_appended_driver, delay, gevent_monkey_patched, is_greenthread_patched, @@ -379,6 +387,9 @@ def test_metadata(self): metadata["driver"]["name"] = "PyMongo|c" else: metadata["driver"]["name"] = "PyMongo" + metadata["driver"]["version"] = _driver_version( + _METADATA["driver"]["version"], metadata["driver"]["name"] + ) metadata["application"] = {"name": "foobar"} client = self.simple_client("mongodb://foo:27017/?appname=foobar&connect=false") options = client.options @@ -390,6 +401,8 @@ def test_metadata(self): self.simple_client(appname="x" * 128) with self.assertRaises(ValueError): self.simple_client(appname="x" * 129) + + def test_metadata_bad_driver_options(self): # Bad "driver" options. self.assertRaises(TypeError, DriverInfo, "Foo", 1, "a") self.assertRaises(TypeError, DriverInfo, version="1", platform="a") @@ -400,12 +413,9 @@ def test_metadata(self): self.simple_client(driver="abc") with self.assertRaises(TypeError): self.simple_client(driver=("Foo", "1", "a")) - # Test appending to driver info. - if has_c(): - metadata["driver"]["name"] = "PyMongo|c|FooDriver" - else: - metadata["driver"]["name"] = "PyMongo|FooDriver" - metadata["driver"]["version"] = "{}|1.2.3".format(_METADATA["driver"]["version"]) + + def test_metadata_appends_driver_info(self): + metadata = _metadata_with_appended_driver(_IS_SYNC, "FooDriver", "1.2.3") client = self.simple_client( "foo", 27017, @@ -413,9 +423,9 @@ def test_metadata(self): driver=DriverInfo("FooDriver", "1.2.3", None), connect=False, ) - options = client.options - self.assertEqual(options.pool_options.metadata, metadata) - metadata["platform"] = "{}|FooPlatform".format(_METADATA["platform"]) + self.assertEqual(client.options.pool_options.metadata, metadata) + + metadata = _metadata_with_appended_driver(_IS_SYNC, "FooDriver", "1.2.3", "FooPlatform") client = self.simple_client( "foo", 27017, @@ -423,27 +433,131 @@ def test_metadata(self): driver=DriverInfo("FooDriver", "1.2.3", "FooPlatform"), connect=False, ) - options = client.options - self.assertEqual(options.pool_options.metadata, metadata) - # Test truncating driver info metadata. + self.assertEqual(client.options.pool_options.metadata, metadata) + + def test_metadata_truncates_driver_info(self): + # Truncated driver info must stay within the limit and keep name and + # version index-aligned. client = self.simple_client( driver=DriverInfo(name="s" * _MAX_METADATA_SIZE), connect=False, ) options = client.options + truncated = options.pool_options.metadata["driver"] self.assertLessEqual( len(bson.encode(options.pool_options.metadata)), _MAX_METADATA_SIZE, ) + self.assertEqual( + truncated["name"].count("|"), + truncated["version"].count("|"), + ) client = self.simple_client( driver=DriverInfo(name="s" * _MAX_METADATA_SIZE, version="s" * _MAX_METADATA_SIZE), connect=False, ) options = client.options + truncated = options.pool_options.metadata["driver"] self.assertLessEqual( len(bson.encode(options.pool_options.metadata)), _MAX_METADATA_SIZE, ) + self.assertEqual( + truncated["name"].count("|"), + truncated["version"].count("|"), + ) + # An oversized wrapper name with no version must retain a truncated + # name rather than collapse to the base entry. + client = self.simple_client( + driver=DriverInfo(name="x" * (_MAX_METADATA_SIZE * 2), version=None), + connect=False, + ) + truncated = client.options.pool_options.metadata["driver"] + self.assertLessEqual( + len(bson.encode(client.options.pool_options.metadata)), + _MAX_METADATA_SIZE, + ) + self.assertIn("xxxx", truncated["name"]) + self.assertEqual( + truncated["name"].count("|"), + truncated["version"].count("|"), + ) + + def test_metadata_truncation_drops_blank_pairs_first(self): + # A platform-only driver appends a blank name/version pair. Truncation + # must drop that pair before trimming a real driver's content. + def metadata_for(name_len: int) -> dict[str, Any]: + name = "PyMongo||" + "W" * name_len # blank pair in the middle + return {"driver": {"name": name, "version": "1.0||1.0"}} + + # Size the name so the document is exactly two bytes over the limit: + # dropping the blank pair alone must make it fit. + name_len = next( + n + for n in range(1, 2 * _MAX_METADATA_SIZE) + if len(bson.encode(metadata_for(n))) == _MAX_METADATA_SIZE + 2 + ) + metadata = metadata_for(name_len) + _truncate_metadata(metadata) + self.assertEqual(len(bson.encode(metadata)), _MAX_METADATA_SIZE) + self.assertEqual( + metadata["driver"], + {"name": "PyMongo|" + "W" * name_len, "version": "1.0|1.0"}, + ) + + def test_metadata_append_is_bounded(self): + # Successive appends must stay within the limit and keep name and + # version index-aligned after truncation. Once the metadata saturates, + # further appends must not grow the dedup tracking list. + client = self.simple_client(connect=False) + for i in range(300): + client.append_metadata(DriverInfo(name=f"D{i}", version=f"1.{i}")) + pool = client.options.pool_options + self.assertLessEqual(len(bson.encode(pool.metadata)), _MAX_METADATA_SIZE) + self.assertEqual( + pool.metadata["driver"]["name"].count("|"), + pool.metadata["driver"]["version"].count("|"), + ) + count = len(pool._PoolOptions__appended_drivers) + for i in range(300, 600): + client.append_metadata(DriverInfo(name=f"D{i}", version=f"1.{i}")) + self.assertEqual(len(pool._PoolOptions__appended_drivers), count) + # Platform-only appends (empty name/version) stay bounded the same way. + client = self.simple_client(connect=False) + for i in range(300): + client.append_metadata(DriverInfo(name="", version="", platform=f"P{i}")) + pool = client.options.pool_options + self.assertLessEqual(len(bson.encode(pool.metadata)), _MAX_METADATA_SIZE) + count = len(pool._PoolOptions__appended_drivers) + for i in range(300, 600): + client.append_metadata(DriverInfo(name="", version="", platform=f"P{i}")) + self.assertEqual(len(pool._PoolOptions__appended_drivers), count) + + def test_metadata_recreates_platform_after_truncation(self): + # Appending a platform after truncation has dropped it recreates the field. + client = self.simple_client(connect=False) + for i in range(300): + client.append_metadata(DriverInfo(name="", version="", platform=f"Q{i}")) + pool = client.options.pool_options + self.assertLess(len(pool._PoolOptions__appended_drivers), 300) + client.append_metadata(DriverInfo(name="Wrapper", version="1.0", platform="Recreated")) + self.assertLessEqual(len(bson.encode(pool.metadata)), _MAX_METADATA_SIZE) + self.assertEqual( + pool.metadata["driver"]["name"].count("|"), + pool.metadata["driver"]["version"].count("|"), + ) + + def test_metadata_deduplicates_none_and_empty(self): + # Empty strings are treated as unset, so a duplicate differing only in + # None vs "" is a no-op. + client = self.simple_client(connect=False) + client.append_metadata(DriverInfo("library", None, "Library Platform")) + names = client.options.pool_options.metadata["driver"]["name"] + vers = client.options.pool_options.metadata["driver"]["version"] + client.append_metadata(DriverInfo("library", "", "Library Platform")) + metadata = client.options.pool_options.metadata + self.assertEqual(metadata["driver"]["name"], names) + self.assertEqual(metadata["driver"]["version"], vers) @mock.patch.dict("os.environ", {ENV_VAR_K8S: "1"}) def test_container_metadata(self): @@ -2177,6 +2291,9 @@ def _test_handshake(self, env_vars, expected_env): metadata["driver"]["name"] = "PyMongo|c" else: metadata["driver"]["name"] = "PyMongo" + metadata["driver"]["version"] = _driver_version( + _METADATA["driver"]["version"], metadata["driver"]["name"] + ) if expected_env is not None: metadata["env"] = expected_env diff --git a/test/test_client_metadata.py b/test/test_client_metadata.py index f5ec92f2f3..4ea3eb00ba 100644 --- a/test/test_client_metadata.py +++ b/test/test_client_metadata.py @@ -18,7 +18,7 @@ import pathlib import time import unittest -from typing import Any, Optional +from typing import Any, Optional, cast import pytest @@ -99,20 +99,14 @@ def check_metadata_added( new_name, new_version, new_platform, new_metadata = self.send_ping_and_get_metadata( client, True ) - if add_name is not None and add_name.lower() in name.lower().split("|"): - self.assertEqual(name, new_name) - self.assertEqual(version, new_version) - self.assertEqual(platform, new_platform) - else: - self.assertEqual(new_name, f"{name}|{add_name}" if add_name is not None else name) - self.assertEqual( - new_version, - f"{version}|{add_version}" if add_version is not None else version, - ) - self.assertEqual( - new_platform, - f"{platform}|{add_platform}" if add_platform is not None else platform, - ) + # Name and version always get a delimiter (empty string if None) to + # preserve 1:1 index correspondence. + self.assertEqual(new_name, f"{name}|{add_name or ''}") + self.assertEqual(new_version, f"{version}|{add_version or ''}") + self.assertEqual( + new_platform, + f"{platform}|{add_platform}" if add_platform is not None else platform, + ) metadata.pop("driver") metadata.pop("platform") @@ -120,7 +114,7 @@ def check_metadata_added( new_metadata.pop("platform") self.assertEqual(metadata, new_metadata) - def test_append_metadata(self): + def test_1_test_that_the_driver_updates_metadata(self): client = self.rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -128,7 +122,7 @@ def test_append_metadata(self): ) self.check_metadata_added(client, "framework", "2.0", "Framework Platform") - def test_append_metadata_platform_none(self): + def test_1_test_that_the_driver_updates_metadata_platform_none(self): client = self.rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -136,7 +130,7 @@ def test_append_metadata_platform_none(self): ) self.check_metadata_added(client, "framework", "2.0", None) - def test_append_metadata_version_none(self): + def test_1_test_that_the_driver_updates_metadata_version_none(self): client = self.rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -144,7 +138,7 @@ def test_append_metadata_version_none(self): ) self.check_metadata_added(client, "framework", None, "Framework Platform") - def test_append_metadata_platform_version_none(self): + def test_1_test_that_the_driver_updates_metadata_platform_version_none(self): client = self.rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -152,14 +146,14 @@ def test_append_metadata_platform_version_none(self): ) self.check_metadata_added(client, "framework", None, None) - def test_multiple_successive_metadata_updates(self): + def test_2_multiple_successive_metadata_updates(self): client = self.rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, connect=False ) client.append_metadata(DriverInfo("library", "1.2", "Library Platform")) self.check_metadata_added(client, "framework", "2.0", "Framework Platform") - def test_multiple_successive_metadata_updates_platform_none(self): + def test_2_multiple_successive_metadata_updates_platform_none(self): client = self.rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -167,7 +161,7 @@ def test_multiple_successive_metadata_updates_platform_none(self): client.append_metadata(DriverInfo("library", "1.2", "Library Platform")) self.check_metadata_added(client, "framework", "2.0", None) - def test_multiple_successive_metadata_updates_version_none(self): + def test_2_multiple_successive_metadata_updates_version_none(self): client = self.rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -175,7 +169,7 @@ def test_multiple_successive_metadata_updates_version_none(self): client.append_metadata(DriverInfo("library", "1.2", "Library Platform")) self.check_metadata_added(client, "framework", None, "Framework Platform") - def test_multiple_successive_metadata_updates_platform_version_none(self): + def test_2_multiple_successive_metadata_updates_platform_version_none(self): client = self.rs_or_single_client( "mongodb://" + self.server.address_string, maxIdleTimeMS=1, @@ -216,10 +210,16 @@ def test_duplicate_driver_name_no_op(self): self.check_metadata_added(client, "framework", None, None) # wait for connection to become idle time.sleep(0.005) - # add same metadata again - self.check_metadata_added(client, "Framework", None, None) + # Append the exact same DriverInfo again: no-op. + name, version, platform, _ = self.send_ping_and_get_metadata(client, True) + time.sleep(0.005) + client.append_metadata(DriverInfo("framework", None, None)) + new_name, new_version, new_platform, _ = self.send_ping_and_get_metadata(client, True) + self.assertEqual(new_name, name) + self.assertEqual(new_version, version) + self.assertEqual(new_platform, platform) - def test_handshake_documents_include_backpressure(self): + def test_9_handshake_documents_include_backpressure(self): # Create a `MongoClient` that is configured to record all handshake documents sent to the server as a part of # connection establishment. client = self.rs_or_single_client("mongodb://" + self.server.address_string) @@ -232,6 +232,125 @@ def test_handshake_documents_include_backpressure(self): # the document has a field `backpressure` whose value is `"2"`. self.assertEqual(self.handshake_req["backpressure"], "2") + def test_10_entries_in_driver_name_and_driver_version_correspond_by_index(self): + cases = [ + ("Gap in middle (name)", [(None, None), ("F2", None)], "||F2", "||"), + ("Gap in middle (version)", [("F1", None), ("F2", "2.0")], "|F1|F2", "||2.0"), + ("Trailing delimiter retained", [("F1", None)], "|F1", "|"), + ( + "Equal versions do not collapse", + [("F1", "{driver_version}")], + "|F1", + "|{driver_version}", + ), + ( + "Equal names do not collapse", + [("{driver_name}", "1.0")], + "|{driver_name}", + "|1.0", + ), + ("Duplicates still deduplicate", [("F1", "1.0"), ("F1", "1.0")], "|F1", "|1.0"), + ("All versions absent", [("F1", None), ("F2", None)], "|F1|F2", "||"), + ("All names absent", [(None, "1.0"), (None, "2.0")], "||", "|1.0|2.0"), + ( + "Non-adjacent duplicate", + [("F1", "1.0"), ("F2", "2.0"), ("F1", "1.0")], + "|F1|F2", + "|1.0|2.0", + ), + ( + "Platform-only difference is not a duplicate", + [("F1", "1.0", "P1"), ("F1", "1.0", "P2")], + "|F1|F1", + "|1.0|1.0", + ), + ( + "Wrapper matching the driver's own identity", + [("{driver_name}", "{driver_version}")], + "|{driver_name}", + "|{driver_version}", + ), + ] + for ( + description, + appended, + expected_name_suffix, + expected_version_suffix, + ) in cases: + with self.subTest(description=description): + client = self.rs_or_single_client( + "mongodb://" + self.server.address_string, + maxIdleTimeMS=1, + ) + self.addCleanup(client.close) + # Capture the driver's own name and version from the first handshake. + name0, version0, _, _ = self.send_ping_and_get_metadata(client, True) + time.sleep(0.005) + + self.assertIsNotNone(name0) + self.assertIsNotNone(version0) + version0 = cast(str, version0) + driver_name = name0.split("|")[0] + driver_version = version0.split("|")[0] + + def resolve(value: Optional[str]) -> Optional[str]: + if value is None: + return None + return value.format(driver_name=driver_name, driver_version=driver_version) + + # Append each DriverInfo in order. + for opts in appended: + d_name = resolve(opts[0]) if len(opts) > 0 else None + d_version = resolve(opts[1]) if len(opts) > 1 else None + d_platform = resolve(opts[2]) if len(opts) > 2 else None + client.append_metadata(DriverInfo(d_name or "", d_version, d_platform)) + + # New handshake with the appended metadata. + name1, version1, _, _ = self.send_ping_and_get_metadata(client, True) + + self.assertEqual( + name1, + name0 + + expected_name_suffix.format( + driver_name=driver_name, driver_version=driver_version + ), + ) + self.assertEqual( + version1, + version0 + + expected_version_suffix.format( + driver_name=driver_name, driver_version=driver_version + ), + ) + + def test_11_appending_metadata_containing_the_delimiter_raises_an_error(self): + cases = [ + ("frame|work", "2.0", "Framework Platform"), + ("framework", "2|0", "Framework Platform"), + ("framework", "2.0", "Framework|Platform"), + ] + for name, version, platform in cases: + with self.subTest(name=name, version=version, platform=platform): + client = self.rs_or_single_client( + "mongodb://" + self.server.address_string, + maxIdleTimeMS=1, + driver=DriverInfo("library", "1.2", "Library Platform"), + ) + self.addCleanup(client.close) + # Send initial handshake. + name0, version0, platform0, _metadata = self.send_ping_and_get_metadata( + client, True + ) + time.sleep(0.005) + # Constructing metadata containing the delimiter raises. + with self.assertRaises(ValueError): + DriverInfo(name, version, platform) + # Metadata is unchanged on the next handshake. + name1, version1, platform1, _ = self.send_ping_and_get_metadata(client, True) + self.assertEqual(name1, name0) + self.assertEqual(version1, version0) + self.assertEqual(platform1, platform0) + if __name__ == "__main__": unittest.main() diff --git a/test/utils_shared.py b/test/utils_shared.py index 6ae9405207..9058019e46 100644 --- a/test/utils_shared.py +++ b/test/utils_shared.py @@ -31,9 +31,11 @@ from collections import abc, defaultdict from functools import partial from inspect import iscoroutinefunction +from typing import Any from bson.objectid import ObjectId from pymongo import monitoring, operations, read_preferences +from pymongo.common import has_c from pymongo.cursor_shared import CursorType from pymongo.errors import OperationFailure from pymongo.helpers_shared import _SENSITIVE_COMMANDS @@ -51,6 +53,7 @@ PoolCreatedEvent, PoolReadyEvent, ) +from pymongo.pool_options import _METADATA from pymongo.pool_shared import _CancellationContext, _PoolGeneration from pymongo.read_concern import ReadConcern from pymongo.server_type import SERVER_TYPE @@ -773,3 +776,36 @@ def pack_msg_header(length: int, request_id: int, response_to: int, op_code: int production header-packing never does. """ return struct.pack(" str: + """Build a metadata driver version aligned 1:1 with ``name`` segments. + + The ``|c`` and ``|async`` name segments always have an empty version entry, + so the version string has one delimiter per name delimiter. ``last_version`` + is used when the final segment carries a wrapped driver's version. + """ + segments = [""] * name.count("|") + if last_version is not None: + segments[-1] = last_version + return "|".join([base_version, *segments]) + + +def _metadata_with_appended_driver( + is_sync: bool, name: str, version: str, platform: str | None = None +) -> dict[str, Any]: + """Build the expected client metadata after appending a driver.""" + driver_name = "PyMongo" + if has_c(): + driver_name += "|c" + if not is_sync: + driver_name += "|async" + metadata = copy.deepcopy(_METADATA) + metadata["driver"]["name"] = driver_name + f"|{name}" + metadata["driver"]["version"] = _driver_version( + _METADATA["driver"]["version"], metadata["driver"]["name"], last_version=version + ) + metadata["application"] = {"name": "foobar"} + if platform is not None: + metadata["platform"] = "{}|{}".format(_METADATA["platform"], platform) + return metadata