Skip to content

Commit 41c758c

Browse files
authored
feat: add closing methods for Qdrant (#3671)
1 parent b3b18aa commit 41c758c

7 files changed

Lines changed: 183 additions & 2 deletions

File tree

integrations/qdrant/src/haystack_integrations/components/retrievers/qdrant/retriever.py

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -130,6 +130,18 @@ def from_dict(cls, data: dict[str, Any]) -> "QdrantEmbeddingRetriever":
130130
data["init_parameters"]["filter_policy"] = FilterPolicy.from_str(filter_policy)
131131
return default_from_dict(cls, data)
132132

133+
def close(self) -> None:
134+
"""
135+
Release the synchronous resources of the underlying Document Store.
136+
"""
137+
self._document_store.close()
138+
139+
async def close_async(self) -> None:
140+
"""
141+
Release the asynchronous resources of the underlying Document Store.
142+
"""
143+
await self._document_store.close_async()
144+
133145
@component.output_types(documents=list[Document])
134146
def run(
135147
self,
@@ -358,6 +370,18 @@ def from_dict(cls, data: dict[str, Any]) -> "QdrantSparseEmbeddingRetriever":
358370
data["init_parameters"]["filter_policy"] = FilterPolicy.from_str(filter_policy)
359371
return default_from_dict(cls, data)
360372

373+
def close(self) -> None:
374+
"""
375+
Release the synchronous resources of the underlying Document Store.
376+
"""
377+
self._document_store.close()
378+
379+
async def close_async(self) -> None:
380+
"""
381+
Release the asynchronous resources of the underlying Document Store.
382+
"""
383+
await self._document_store.close_async()
384+
361385
@component.output_types(documents=list[Document])
362386
def run(
363387
self,
@@ -596,6 +620,18 @@ def from_dict(cls, data: dict[str, Any]) -> "QdrantHybridRetriever":
596620
data["init_parameters"]["filter_policy"] = FilterPolicy.from_str(filter_policy)
597621
return default_from_dict(cls, data)
598622

623+
def close(self) -> None:
624+
"""
625+
Release the synchronous resources of the underlying Document Store.
626+
"""
627+
self._document_store.close()
628+
629+
async def close_async(self) -> None:
630+
"""
631+
Release the asynchronous resources of the underlying Document Store.
632+
"""
633+
await self._document_store.close_async()
634+
599635
@component.output_types(documents=list[Document])
600636
def run(
601637
self,

integrations/qdrant/src/haystack_integrations/document_stores/qdrant/document_store.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import inspect
22
from collections.abc import AsyncGenerator, Generator
3+
from contextlib import suppress
34
from dataclasses import replace
45
from itertools import islice
56
from typing import Any, ClassVar, cast
@@ -301,6 +302,24 @@ async def _initialize_async_client(self) -> None:
301302
self.payload_fields_to_index,
302303
)
303304

305+
def close(self) -> None:
306+
"""
307+
Release the associated synchronous resources.
308+
"""
309+
if self._client is not None:
310+
with suppress(Exception):
311+
self._client.close()
312+
self._client = None
313+
314+
async def close_async(self) -> None:
315+
"""
316+
Release the associated asynchronous resources.
317+
"""
318+
if self._async_client is not None:
319+
with suppress(Exception):
320+
await self._async_client.close()
321+
self._async_client = None
322+
304323
def count_documents(self) -> int:
305324
"""
306325
Returns the number of documents present in the Document Store.

integrations/qdrant/tests/test_document_store.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -370,6 +370,29 @@ def test_metadata_methods_swallow_client_errors(self, method_name, args, expecte
370370
):
371371
assert getattr(document_store, method_name)(*args) == expected
372372

373+
def test_close(self):
374+
document_store = QdrantDocumentStore(location=":memory:")
375+
mock_client = MagicMock()
376+
document_store._client = mock_client
377+
378+
document_store.close()
379+
380+
mock_client.close.assert_called_once()
381+
assert document_store._client is None
382+
383+
document_store.close()
384+
mock_client.close.assert_called_once()
385+
386+
def test_close_is_exception_safe(self):
387+
document_store = QdrantDocumentStore(location=":memory:")
388+
mock_client = MagicMock()
389+
mock_client.close.side_effect = RuntimeError("boom")
390+
document_store._client = mock_client
391+
392+
document_store.close()
393+
394+
assert document_store._client is None
395+
373396

374397
@pytest.mark.integration
375398
class TestQdrantDocumentStore(
@@ -409,6 +432,16 @@ def assert_documents_are_equal(self, received: list[Document], expected: list[Do
409432
# Check that the sets are equal, meaning the content and IDs match regardless of order
410433
assert {doc.id for doc in received} == {doc.id for doc in expected}
411434

435+
def test_close_and_reopen(self, document_store: QdrantDocumentStore):
436+
assert document_store.count_documents() == 0
437+
assert document_store._client is not None
438+
439+
document_store.close()
440+
441+
assert document_store._client is None
442+
assert document_store.count_documents() == 0
443+
assert document_store._client is not None
444+
412445
def test_prepare_client_params_no_mutability(self):
413446
metadata = {"key": "value"}
414447
doc_store = QdrantDocumentStore(

integrations/qdrant/tests/test_document_store_async.py

Lines changed: 34 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
from unittest.mock import MagicMock, patch
1+
from unittest.mock import AsyncMock, MagicMock, patch
22

33
import pytest
44
import pytest_asyncio
@@ -162,6 +162,29 @@ async def test_metadata_methods_async_absorb_client_errors(self, method_name, ar
162162
):
163163
assert await getattr(document_store, method_name)(*args) == expected
164164

165+
async def test_close_async(self):
166+
document_store = QdrantDocumentStore(location=":memory:")
167+
mock_client = AsyncMock()
168+
document_store._async_client = mock_client
169+
170+
await document_store.close_async()
171+
172+
mock_client.close.assert_awaited_once()
173+
assert document_store._async_client is None
174+
175+
await document_store.close_async()
176+
mock_client.close.assert_awaited_once()
177+
178+
async def test_close_async_is_exception_safe(self):
179+
document_store = QdrantDocumentStore(location=":memory:")
180+
mock_client = AsyncMock()
181+
mock_client.close.side_effect = RuntimeError("boom")
182+
document_store._async_client = mock_client
183+
184+
await document_store.close_async()
185+
186+
assert document_store._async_client is None
187+
165188

166189
@pytest.mark.integration
167190
@pytest.mark.asyncio
@@ -195,6 +218,16 @@ def assert_documents_are_equal(self, received: list[Document], expected: list[Do
195218
assert len(received) == len(expected)
196219
assert {doc.id for doc in received} == {doc.id for doc in expected}
197220

221+
async def test_close_async_and_reopen(self, document_store: QdrantDocumentStore):
222+
assert await document_store.count_documents_async() == 0
223+
assert document_store._async_client is not None
224+
225+
await document_store.close_async()
226+
227+
assert document_store._async_client is None
228+
assert await document_store.count_documents_async() == 0
229+
assert document_store._async_client is not None
230+
198231
@pytest.mark.asyncio
199232
async def test_write_documents_async(self, document_store: QdrantDocumentStore):
200233
docs = [Document(id="1")]

integrations/qdrant/tests/test_embedding_retriever.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -146,6 +146,26 @@ def test_from_dict_no_filter_policy(self):
146146
retriever = QdrantEmbeddingRetriever.from_dict(data)
147147
assert retriever._filter_policy == FilterPolicy.REPLACE # defaults to REPLACE
148148

149+
def test_close(self):
150+
mock_store = Mock(spec=QdrantDocumentStore)
151+
retriever = QdrantEmbeddingRetriever(document_store=mock_store)
152+
153+
retriever.close()
154+
155+
mock_store.close.assert_called_once_with()
156+
assert retriever._document_store is mock_store
157+
158+
@pytest.mark.asyncio
159+
async def test_close_async(self):
160+
mock_store = Mock(spec=QdrantDocumentStore)
161+
mock_store.close_async = AsyncMock()
162+
retriever = QdrantEmbeddingRetriever(document_store=mock_store)
163+
164+
await retriever.close_async()
165+
166+
mock_store.close_async.assert_awaited_once_with()
167+
assert retriever._document_store is mock_store
168+
149169
def test_run(self):
150170
mock_store = Mock(spec=QdrantDocumentStore)
151171
mock_store._query_by_embedding.return_value = [Document(content="doc", embedding=[0.1, 0.2])]

integrations/qdrant/tests/test_hybrid_retriever.py

Lines changed: 21 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
from unittest.mock import Mock
1+
from unittest.mock import AsyncMock, Mock
22

33
import pytest
44
from haystack.dataclasses import Document, SparseEmbedding
@@ -146,6 +146,26 @@ def test_from_dict_no_filter_policy(self):
146146
assert retriever._group_by is None
147147
assert retriever._group_size is None
148148

149+
def test_close(self):
150+
mock_store = Mock(spec=QdrantDocumentStore)
151+
retriever = QdrantHybridRetriever(document_store=mock_store)
152+
153+
retriever.close()
154+
155+
mock_store.close.assert_called_once_with()
156+
assert retriever._document_store is mock_store
157+
158+
@pytest.mark.asyncio
159+
async def test_close_async(self):
160+
mock_store = Mock(spec=QdrantDocumentStore)
161+
mock_store.close_async = AsyncMock()
162+
retriever = QdrantHybridRetriever(document_store=mock_store)
163+
164+
await retriever.close_async()
165+
166+
mock_store.close_async.assert_awaited_once_with()
167+
assert retriever._document_store is mock_store
168+
149169
def test_run(self):
150170
mock_store = Mock(spec=QdrantDocumentStore)
151171
sparse_embedding = SparseEmbedding(indices=[0, 1, 2, 3], values=[0.1, 0.8, 0.05, 0.33])

integrations/qdrant/tests/test_sparse_embedding_retriever.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -154,6 +154,26 @@ def test_from_dict_no_filter_policy(self):
154154
assert retriever._group_by is None
155155
assert retriever._group_size is None
156156

157+
def test_close(self):
158+
mock_store = Mock(spec=QdrantDocumentStore)
159+
retriever = QdrantSparseEmbeddingRetriever(document_store=mock_store)
160+
161+
retriever.close()
162+
163+
mock_store.close.assert_called_once_with()
164+
assert retriever._document_store is mock_store
165+
166+
@pytest.mark.asyncio
167+
async def test_close_async(self):
168+
mock_store = Mock(spec=QdrantDocumentStore)
169+
mock_store.close_async = AsyncMock()
170+
retriever = QdrantSparseEmbeddingRetriever(document_store=mock_store)
171+
172+
await retriever.close_async()
173+
174+
mock_store.close_async.assert_awaited_once_with()
175+
assert retriever._document_store is mock_store
176+
157177
def test_run(self):
158178
mock_store = Mock(spec=QdrantDocumentStore)
159179
sparse = SparseEmbedding(indices=[0, 5], values=[0.1, 0.7])

0 commit comments

Comments
 (0)