diff --git a/doc/changelog.rst b/doc/changelog.rst index 41e6dee233..60ff019b4b 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -10,8 +10,15 @@ 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 + (`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..54905e8872 100644 --- a/pymongo/driver_info.py +++ b/pymongo/driver_info.py @@ -31,6 +31,10 @@ class DriverInfo(namedtuple("DriverInfo", ["name", "version", "platform"])): 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. + + 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 +46,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..b9a007c738 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,6 +202,25 @@ 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 "", + ) + + # 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.""" @@ -226,7 +247,7 @@ def _truncate_metadata(metadata: MutableMapping[str, Any]) -> None: overflow = encoded_size - _MAX_METADATA_SIZE plat = metadata.get("platform", "") if plat: - plat = plat[:-overflow] + plat = _truncate_utf8(plat, overflow) if plat: metadata["platform"] = plat else: @@ -234,26 +255,36 @@ def _truncate_metadata(metadata: MutableMapping[str, Any]) -> None: encoded_size = len(bson.encode(metadata)) if encoded_size <= _MAX_METADATA_SIZE: return - # 5. Truncate driver info. - overflow = encoded_size - _MAX_METADATA_SIZE + # 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"] + # Trim wrapper version and name content first, dropping paired segments + # only as a last resort, so name and version stay 1:1 aligned. + while True: + encoded_size = len(bson.encode(metadata)) + if encoded_size <= _MAX_METADATA_SIZE: + break + overflow = encoded_size - _MAX_METADATA_SIZE + previous = (driver.get("name", ""), driver.get("version", "")) + n_parts = driver.get("name", "").split("|") + v_parts = driver.get("version", "").split("|") + + if len(v_parts) > 1 and v_parts[-1]: + v_parts[-1] = _truncate_utf8(v_parts[-1], overflow) + driver["version"] = "|".join(v_parts) + elif len(n_parts) > 1 and n_parts[-1]: + n_parts[-1] = _truncate_utf8(n_parts[-1], overflow) + driver["name"] = "|".join(n_parts) + elif len(n_parts) > 1: + n_parts.pop() + v_parts.pop() + driver["name"] = "|".join(n_parts) + driver["version"] = "|".join(v_parts) + else: + break + + if previous == (driver.get("name"), driver.get("version")): + break # If the first getaddrinfo call of this interpreter's life is on a thread, @@ -277,6 +308,7 @@ class PoolOptions: """ __slots__ = ( + "__appended_drivers", "__appname", "__compression_settings", "__connect_timeout", @@ -288,6 +320,7 @@ class PoolOptions: "__max_idle_time_seconds", "__max_pool_size", "__metadata", + "__metadata_lock", "__min_pool_size", "__pause_enabled", "__server_api", @@ -336,6 +369,11 @@ def __init__( self.__load_balanced = load_balanced self.__credentials = credentials self.__metadata = copy.deepcopy(_METADATA) + self.__appended_drivers: list[DriverInfo] = [] + # 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 +391,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 +414,38 @@ 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("|") + metadata = copy.deepcopy(self.__metadata) + + 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 and the dedup membership check stays fast. + if ( + metadata["driver"]["name"].count("|") > name_delims + and metadata["driver"]["version"].count("|") > version_delims + ): + self.__appended_drivers.append(driver) @property def _credentials(self) -> Optional[MongoCredential]: diff --git a/test/asynchronous/test_client.py b/test/asynchronous/test_client.py index 90a2d33a45..b3b5a7e648 100644 --- a/test/asynchronous/test_client.py +++ b/test/asynchronous/test_client.py @@ -124,6 +124,7 @@ NTHREADS, CMAPListener, FunctionCallRecorder, + _driver_version, delay, gevent_monkey_patched, is_greenthread_patched, @@ -380,12 +381,31 @@ async def test_read_preference(self): ) self.assertEqual(c.read_preference, ReadPreference.NEAREST) + def _metadata_with_appended_driver( + self, name: str, version: str, platform: str | None = None + ) -> dict[str, Any]: + metadata = copy.deepcopy(_METADATA) + if has_c(): + metadata["driver"]["name"] = "PyMongo|c|async|" + name + else: + metadata["driver"]["name"] = "PyMongo|async|" + 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 + async def test_metadata(self): metadata = copy.deepcopy(_METADATA) if has_c(): 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 +417,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 +429,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 = self._metadata_with_appended_driver("FooDriver", "1.2.3") client = self.simple_client( "foo", 27017, @@ -420,9 +439,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 = self._metadata_with_appended_driver("FooDriver", "1.2.3", "FooPlatform") client = self.simple_client( "foo", 27017, @@ -430,27 +449,116 @@ 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_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_rejects_delimiter(self): + # The '|' delimiter is reserved for joining appended metadata, so it + # must be rejected in every field. + self.assertRaises(ValueError, DriverInfo, "a|b", "1.0", None) + self.assertRaises(ValueError, DriverInfo, "lib", "1|0", None) + self.assertRaises(ValueError, DriverInfo, "lib", "1.0", "Frame|Platform") + + 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 +2332,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..7a6833753c 100644 --- a/test/test_client.py +++ b/test/test_client.py @@ -123,6 +123,7 @@ NTHREADS, CMAPListener, FunctionCallRecorder, + _driver_version, delay, gevent_monkey_patched, is_greenthread_patched, @@ -373,12 +374,31 @@ def test_read_preference(self): ) self.assertEqual(c.read_preference, ReadPreference.NEAREST) + def _metadata_with_appended_driver( + self, name: str, version: str, platform: str | None = None + ) -> dict[str, Any]: + metadata = copy.deepcopy(_METADATA) + if has_c(): + metadata["driver"]["name"] = "PyMongo|c|" + name + else: + metadata["driver"]["name"] = "PyMongo|" + 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 + def test_metadata(self): metadata = copy.deepcopy(_METADATA) if has_c(): 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 +410,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 +422,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 = self._metadata_with_appended_driver("FooDriver", "1.2.3") client = self.simple_client( "foo", 27017, @@ -413,9 +432,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 = self._metadata_with_appended_driver("FooDriver", "1.2.3", "FooPlatform") client = self.simple_client( "foo", 27017, @@ -423,27 +442,116 @@ 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_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_rejects_delimiter(self): + # The '|' delimiter is reserved for joining appended metadata, so it + # must be rejected in every field. + self.assertRaises(ValueError, DriverInfo, "a|b", "1.0", None) + self.assertRaises(ValueError, DriverInfo, "lib", "1|0", None) + self.assertRaises(ValueError, DriverInfo, "lib", "1.0", "Frame|Platform") + + 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 +2285,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..341e9c90dd 100644 --- a/test/utils_shared.py +++ b/test/utils_shared.py @@ -773,3 +773,16 @@ 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])