diff --git a/litellm/llms/mongodb/common_utils.py b/litellm/llms/mongodb/common_utils.py index 12391b80959..4aafaf86a5e 100644 --- a/litellm/llms/mongodb/common_utils.py +++ b/litellm/llms/mongodb/common_utils.py @@ -9,6 +9,8 @@ TLS handshake and topology discovery: measured at ~890ms against Atlas versus """ import asyncio +import weakref +from asyncio import AbstractEventLoop from dataclasses import dataclass from typing import TYPE_CHECKING, Final @@ -53,7 +55,12 @@ class MongoClientKey: _sync_clients: dict[MongoClientKey, "MongoClient"] = {} # mutable-ok: process-level connection cache, see module docstring -_async_clients: dict[tuple[MongoClientKey, int], "AsyncMongoClient"] = {} # mutable-ok: same cache, keyed per event loop +# The value carries a weak reference to the loop the client was built on: CPython recycles +# id() aggressively (measured: 200 of 200 fresh loops landed on an id already in this cache), +# so the id alone would hand a new loop a client bound to a closed one. +_async_clients: dict[ # mutable-ok: same cache, keyed per event loop + tuple[MongoClientKey, int], tuple["weakref.ref[AbstractEventLoop]", "AsyncMongoClient"] +] = {} def import_sync_mongo_client() -> "type[MongoClient]": @@ -93,13 +100,14 @@ def get_sync_client(key: MongoClientKey) -> "MongoClient": def get_async_client(key: MongoClientKey) -> "AsyncMongoClient": """Async clients bind to the loop that created them, so the cache is keyed per loop.""" - loop_key: Final = (key, id(asyncio.get_running_loop())) + loop: Final = asyncio.get_running_loop() + loop_key: Final = (key, id(loop)) cached: Final = _async_clients.get(loop_key) - if cached is not None: - return cached + if cached is not None and cached[0]() is loop: + return cached[1] client: Final = import_async_mongo_client()(key.connection_string, **_client_kwargs(key)) - if len(_async_clients) < _MAX_CACHED_CLIENTS: - _async_clients[loop_key] = client + if len(_async_clients) < _MAX_CACHED_CLIENTS or loop_key in _async_clients: + _async_clients[loop_key] = (weakref.ref(loop), client) return client diff --git a/tests/test_litellm/llms/mongodb/vector_stores/test_mongodb_transformation.py b/tests/test_litellm/llms/mongodb/vector_stores/test_mongodb_transformation.py index 15e1b1efab5..9e2bccc3f26 100644 --- a/tests/test_litellm/llms/mongodb/vector_stores/test_mongodb_transformation.py +++ b/tests/test_litellm/llms/mongodb/vector_stores/test_mongodb_transformation.py @@ -1,4 +1,7 @@ +import asyncio +import gc import sys +import weakref from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -541,6 +544,46 @@ class TestClientCache: assert first is second + def test_a_new_loop_never_inherits_a_closed_loop_client(self): + """CPython recycles id() so aggressively that a fresh event loop almost always lands on + the id of one already collected: measured at 37 of 40 rounds. Keying the cache on the id + alone therefore hands the new loop an AsyncMongoClient bound to a closed loop, and every + operation on it raises "Event loop is closed".""" + + class LoopAgnosticClient: + """Holds no reference to the loop, unlike pymongo's, whose own reference happens to + keep ids from being recycled and hides the bug until the cache fills.""" + + def __init__(self, *args, **kwargs): + self.built_on = None + + key = self._key() + clients_handed_out = [] + + async def fetch(): + return get_async_client(key) + + with patch("litellm.llms.mongodb.common_utils.import_async_mongo_client") as importer: + importer.return_value = LoopAgnosticClient + + for _ in range(20): + loop = asyncio.new_event_loop() + client = loop.run_until_complete(fetch()) + clients_handed_out.append((client, client.built_on, loop.is_closed())) + client.built_on = weakref.ref(loop) + loop.close() + del loop + gc.collect() + + stale = [ + handed_out + for client, built_on, _ in clients_handed_out + if built_on is not None and (built_on() is None or built_on().is_closed()) + for handed_out in (client,) + ] + assert stale == [], f"{len(stale)} of 20 loops were handed a client built on a closed loop" + + class TestClientKeyDerivation: def test_no_timeout_uses_the_bounded_defaults(self): key = MongoDBVectorStoreConfig._client_key(_MongoDBSearchParams.model_validate(BASE_PARAMS), None)