From 2a11c2747f58f24a1c9f1babc30027afa9ec2a8a Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 4 Sep 2026 15:25:38 -0700 Subject: [PATCH] fix(vector_stores): serialize the MongoDB client cache so concurrent searches cannot trip over an eviction Async searches reach the sync client through executor threads, so the LRU cache is shared state. A key could be evicted between the lookup and the reordering that followed it, and the reordering then raised KeyError and became a 500. Reproduced at 15 failures per run with 16 threads over 34 keys and a 1ns switch interval; the regression test is that workload. --- litellm/llms/mongodb/common_utils.py | 40 ++++++++++++------- .../test_mongodb_transformation.py | 27 +++++++++++++ 2 files changed, 53 insertions(+), 14 deletions(-) diff --git a/litellm/llms/mongodb/common_utils.py b/litellm/llms/mongodb/common_utils.py index 27a2a96bd1f..02c0b359407 100644 --- a/litellm/llms/mongodb/common_utils.py +++ b/litellm/llms/mongodb/common_utils.py @@ -2,6 +2,7 @@ so every import of it is deferred to call time.""" import asyncio +import threading import weakref from asyncio import AbstractEventLoop from collections import OrderedDict @@ -69,14 +70,23 @@ _AsyncClientCache: TypeAlias = "OrderedDict[_AsyncClientCacheKey, _AsyncClientEn _sync_clients: Final[_SyncClientCache] = OrderedDict() # mutable-ok: process-level client cache _async_clients: Final[_AsyncClientCache] = OrderedDict() # mutable-ok: same cache, per loop +# async searches reach the sync client through executor threads, so both caches are shared state +_cache_lock: Final = threading.Lock() def _store_bounded(cache: "OrderedDict[_K, _V]", cache_key: "_K", value: "_V") -> None: """Eviction only drops this cache's reference; an in-flight search keeps its client alive.""" - cache[cache_key] = value # mutable-ok: an LRU cache is mutable state by definition - cache.move_to_end(cache_key) - while len(cache) > _MAX_CACHED_CLIENTS: - cache.popitem(last=False) + with _cache_lock: + cache[cache_key] = value # mutable-ok: an LRU cache is mutable state by definition + cache.move_to_end(cache_key) + while len(cache) > _MAX_CACHED_CLIENTS: + cache.popitem(last=False) + + +def _mark_used(cache: "OrderedDict[_K, _V]", cache_key: "_K") -> None: + with _cache_lock: + if cache_key in cache: + cache.move_to_end(cache_key) def import_sync_mongo_client() -> "type[MongoClient]": @@ -109,7 +119,7 @@ def _client_kwargs(key: MongoClientKey) -> Mapping[str, object]: def get_sync_client(key: MongoClientKey, client_class: SyncClientFactory | None = None) -> "MongoClient": cached: Final = _sync_clients.get(key) if cached is not None: - _sync_clients.move_to_end(key) + _mark_used(_sync_clients, key) return cached build: Final = client_class if client_class is not None else import_sync_mongo_client() client: Final = build(key.connection_string, **_client_kwargs(key)) @@ -120,12 +130,13 @@ def get_sync_client(key: MongoClientKey, client_class: SyncClientFactory | None def _purge_dead_loops() -> None: """A cached client holds its loop alive, so a closed loop's entry would pin that client and its sockets for the life of the process.""" - for stale in tuple( - cache_key - for cache_key, (loop_ref, _) in _async_clients.items() - if (cached_loop := loop_ref()) is None or cached_loop.is_closed() - ): - del _async_clients[stale] + with _cache_lock: + for stale in tuple( + cache_key + for cache_key, (loop_ref, _) in _async_clients.items() + if (cached_loop := loop_ref()) is None or cached_loop.is_closed() + ): + del _async_clients[stale] def get_async_client(key: MongoClientKey, client_class: AsyncClientFactory | None = None) -> "AsyncMongoClient": @@ -134,7 +145,7 @@ def get_async_client(key: MongoClientKey, client_class: AsyncClientFactory | Non loop_key: Final = (key, id(loop)) cached: Final = _async_clients.get(loop_key) if cached is not None and cached[0]() is loop: - _async_clients.move_to_end(loop_key) + _mark_used(_async_clients, loop_key) return cached[1] _purge_dead_loops() build: Final = client_class if client_class is not None else import_async_mongo_client() @@ -144,8 +155,9 @@ def get_async_client(key: MongoClientKey, client_class: AsyncClientFactory | Non def reset_client_cache() -> None: - _sync_clients.clear() - _async_clients.clear() + with _cache_lock: + _sync_clients.clear() + _async_clients.clear() _AUTHENTICATION_FAILED_CODE: Final = 18 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 faf20f87ae5..f5d31c0da54 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,6 +1,7 @@ import asyncio import gc import sys +import threading import weakref from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -670,6 +671,32 @@ class TestClientCache: assert get_sync_client(newest, RecordingClient) is kept assert oldest not in _sync_clients + def test_concurrent_searches_never_trip_over_an_eviction(self): + """Async searches run the sync client through executor threads, so a key can be evicted + between the lookup and the reordering that follows it.""" + errors = [] + churn = _MAX_CACHED_CLIENTS + 2 + + def hammer(offset): + try: + for step in range(3_000): + get_sync_client(self._key(f"mongodb://h-{(step + offset) % churn}:27017"), RecordingClient) + except Exception as e: + errors.append(repr(e)) + + previous = sys.getswitchinterval() + sys.setswitchinterval(1e-9) + try: + threads = [threading.Thread(target=hammer, args=(offset,)) for offset in range(16)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + finally: + sys.setswitchinterval(previous) + + assert errors == [] + def test_the_cache_never_grows_past_its_cap(self): for slot in range(_MAX_CACHED_CLIENTS * 3): get_sync_client(self._key(f"mongodb://host-{slot}:27017"), RecordingClient)