mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vector_stores): stop MongoDB handing a new event loop a closed loop's client
The async client cache was keyed on id(loop). CPython recycles those ids so aggressively that a fresh event loop nearly always lands on the id of one already collected, measured at 37 of 40 rounds, so the cache handed the new loop an AsyncMongoClient bound to a closed loop and every operation on it raised "Event loop is closed". The entry now carries a weak reference to the loop it was built on and a hit only counts when that reference still points at the running loop, so a recycled id misses and builds a fresh client. A stale entry can also be replaced once the cache is full, which the old size check prevented. pymongo's own client keeps its loop alive, which is why the sync proxy path never saw this; a script calling asyncio.run() per search, or a test suite with a loop per test, does.
This commit is contained in:
parent
1f8cbee8aa
commit
a472484291
2 changed files with 57 additions and 6 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue