Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions docarray/index/backends/mongodb_atlas.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import logging
from dataclasses import dataclass, field
from functools import cached_property
from importlib.metadata import version
from typing import (
Any,
Dict,
Expand All @@ -20,6 +21,7 @@
import bson
import numpy as np
from pymongo import MongoClient
from pymongo.driver_info import DriverInfo

from docarray import BaseDoc, DocList, handler
from docarray.index.abstract import BaseDocIndex, _raise_not_composable
Expand Down Expand Up @@ -115,7 +117,7 @@ def _connect_to_mongodb_atlas(atlas_connection_uri: str):

client = MongoClient(
atlas_connection_uri,
# driver=DriverInfo(name="docarray", version=version("docarray"))
driver=DriverInfo(name="DocArray", version=version("docarray")),
)
return client

Expand Down Expand Up @@ -557,7 +559,6 @@ def _vector_search_stage(
limit: int,
filters: List[Dict[str, Any]] = None,
) -> Dict[str, Any]:

search_index_name = self._get_column_db_index(search_field)
oversampling_factor = self._get_oversampling_factor(search_field)
max_candidates = self._get_max_candidates(search_field)
Expand Down
17 changes: 17 additions & 0 deletions tests/index/mongo_atlas/test_driver_info.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
from unittest.mock import MagicMock, patch

from pymongo.driver_info import DriverInfo


def test_connect_passes_driver_info():
"""_connect_to_mongodb_atlas passes DriverInfo(name='DocArray') to MongoClient."""
from docarray.index.backends.mongodb_atlas import MongoDBAtlasDocumentIndex

with patch("docarray.index.backends.mongodb_atlas.MongoClient") as mock_client_cls:
mock_client_cls.return_value = MagicMock()
MongoDBAtlasDocumentIndex._connect_to_mongodb_atlas("mongodb://localhost")

mock_client_cls.assert_called_once()
_, kwargs = mock_client_cls.call_args
assert isinstance(kwargs.get("driver"), DriverInfo)
assert kwargs["driver"].name == "DocArray"
Loading