diff --git a/haystack/components/caching/cache_checker.py b/haystack/components/caching/cache_checker.py index 15a79667766..2da90c84786 100644 --- a/haystack/components/caching/cache_checker.py +++ b/haystack/components/caching/cache_checker.py @@ -121,3 +121,17 @@ async def run_async(self, items: list[Any]) -> dict[str, Any]: else: misses.append(item) return {"hits": found_documents, "misses": misses} + + def close(self) -> None: + """ + Release the synchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close"): + self.document_store.close() + + async def close_async(self) -> None: + """ + Release the asynchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close_async"): + await self.document_store.close_async() diff --git a/haystack/components/retrievers/auto_merging_retriever.py b/haystack/components/retrievers/auto_merging_retriever.py index 56f29bdcff1..aab97e38421 100644 --- a/haystack/components/retrievers/auto_merging_retriever.py +++ b/haystack/components/retrievers/auto_merging_retriever.py @@ -224,3 +224,17 @@ async def _try_merge_level(docs_to_merge: list[Document], docs_to_return: list[D return await _try_merge_level(merged_docs, docs_to_return) return {"documents": await _try_merge_level(documents, [])} + + def close(self) -> None: + """ + Release the synchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close"): + self.document_store.close() + + async def close_async(self) -> None: + """ + Release the asynchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close_async"): + await self.document_store.close_async() diff --git a/haystack/components/retrievers/filter_retriever.py b/haystack/components/retrievers/filter_retriever.py index a7f893d7452..29c68f11ec5 100644 --- a/haystack/components/retrievers/filter_retriever.py +++ b/haystack/components/retrievers/filter_retriever.py @@ -102,3 +102,17 @@ async def run_async(self, filters: dict[str, Any] | None = None) -> dict[str, An # 'ignore' since filter_documents_async is not defined in the Protocol but exists in the implementations out_documents = await self.document_store.filter_documents_async(filters=filters or self.filters) # type: ignore[attr-defined] return {"documents": out_documents} + + def close(self) -> None: + """ + Release the synchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close"): + self.document_store.close() + + async def close_async(self) -> None: + """ + Release the asynchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close_async"): + await self.document_store.close_async() diff --git a/haystack/components/retrievers/sentence_window_retriever.py b/haystack/components/retrievers/sentence_window_retriever.py index a4494f03d8f..e24b2666d6d 100644 --- a/haystack/components/retrievers/sentence_window_retriever.py +++ b/haystack/components/retrievers/sentence_window_retriever.py @@ -319,3 +319,17 @@ def _build_filter_conditions(self, split_id: int, window_size: int, source_ids: *source_id_filters, ] return {"operator": "AND", "conditions": conditions} + + def close(self) -> None: + """ + Release the synchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close"): + self.document_store.close() + + async def close_async(self) -> None: + """ + Release the asynchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close_async"): + await self.document_store.close_async() diff --git a/haystack/components/writers/document_writer.py b/haystack/components/writers/document_writer.py index a8ab4fbe667..8493bfc336c 100644 --- a/haystack/components/writers/document_writer.py +++ b/haystack/components/writers/document_writer.py @@ -125,3 +125,17 @@ async def run_async(self, documents: list[Document], policy: DuplicatePolicy | N documents_written = await self.document_store.write_documents_async(documents=documents, policy=policy) return {"documents_written": documents_written} + + def close(self) -> None: + """ + Release the synchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close"): + self.document_store.close() + + async def close_async(self) -> None: + """ + Release the asynchronous resources of the underlying Document Store. + """ + if hasattr(self.document_store, "close_async"): + await self.document_store.close_async() diff --git a/releasenotes/notes/closing-methods-for-ds-holding-comps-e2dcfc4e5afddd0f.yaml b/releasenotes/notes/closing-methods-for-ds-holding-comps-e2dcfc4e5afddd0f.yaml new file mode 100644 index 00000000000..7f5decf73c2 --- /dev/null +++ b/releasenotes/notes/closing-methods-for-ds-holding-comps-e2dcfc4e5afddd0f.yaml @@ -0,0 +1,7 @@ +--- +features: + - | + Haystack components that use a document store now provide ``close`` and ``close_async`` methods for releasing + resources. These methods are available on: ``AutoMergingRetriever``, ``CacheChecker``, ``DocumentWriter``, + ``FilterRetriever``, and ``SentenceWindowRetriever``. If the underlying Document Store does not implement the + corresponding method, calling ``close`` or ``close_async`` has no effect. diff --git a/test/components/caching/test_url_cache_checker.py b/test/components/caching/test_cache_checker.py similarity index 89% rename from test/components/caching/test_url_cache_checker.py rename to test/components/caching/test_cache_checker.py index f3330c0b62c..8c36e149eda 100644 --- a/test/components/caching/test_url_cache_checker.py +++ b/test/components/caching/test_cache_checker.py @@ -2,7 +2,7 @@ # # SPDX-License-Identifier: Apache-2.0 -from unittest.mock import MagicMock +from unittest.mock import MagicMock, Mock import pytest @@ -92,3 +92,14 @@ def test_filters_syntax(self): checker.run(items=["https://example.com/1"]) valid_filters_syntax = {"field": "url", "operator": "==", "value": "https://example.com/1"} mocked_docstore_class.filter_documents.assert_any_call(filters=valid_filters_syntax) + + def test_close(self): + closable_document_store = Mock(spec=["close"]) + checker = CacheChecker(document_store=closable_document_store, cache_field="url") + checker.close() + closable_document_store.close.assert_called_once_with() + + nonclosable_document_store = Mock(spec=[]) + checker = CacheChecker(document_store=nonclosable_document_store, cache_field="url") + checker.close() + assert nonclosable_document_store.mock_calls == [] diff --git a/test/components/caching/test_cache_checker_async.py b/test/components/caching/test_cache_checker_async.py index 05666abd353..d33307af875 100644 --- a/test/components/caching/test_cache_checker_async.py +++ b/test/components/caching/test_cache_checker_async.py @@ -2,7 +2,7 @@ # # SPDX-License-Identifier: Apache-2.0 -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, Mock import pytest @@ -60,3 +60,16 @@ async def test_run_async_filters_syntax(self): await checker.run_async(items=["https://example.com/1"]) expected_filters = {"field": "url", "operator": "==", "value": "https://example.com/1"} mock_store.filter_documents_async.assert_awaited_once_with(filters=expected_filters) + + @pytest.mark.asyncio + async def test_close_async(self): + closable_document_store = Mock(spec=["close_async"]) + closable_document_store.close_async = AsyncMock() + checker = CacheChecker(document_store=closable_document_store, cache_field="url") + await checker.close_async() + closable_document_store.close_async.assert_awaited_once_with() + + nonclosable_document_store = Mock(spec=[]) + checker = CacheChecker(document_store=nonclosable_document_store, cache_field="url") + await checker.close_async() + assert nonclosable_document_store.mock_calls == [] diff --git a/test/components/retrievers/test_auto_merging_retriever.py b/test/components/retrievers/test_auto_merging_retriever.py index ee5f53ffb98..ba9fb85fbee 100644 --- a/test/components/retrievers/test_auto_merging_retriever.py +++ b/test/components/retrievers/test_auto_merging_retriever.py @@ -2,6 +2,8 @@ # # SPDX-License-Identifier: Apache-2.0 +from unittest.mock import Mock + import pytest from haystack import Document, Pipeline @@ -252,3 +254,14 @@ def test_run_go_up_hierarchy_multiple_levels_hit_root_document(self, in_memory_d assert len(result["documents"]) == 1 assert result["documents"][0].meta["__level"] == 0 # hit root document + + def test_close(self): + closable_document_store = Mock(spec=["close"]) + retriever = AutoMergingRetriever(document_store=closable_document_store) + retriever.close() + closable_document_store.close.assert_called_once_with() + + nonclosable_document_store = Mock(spec=[]) + retriever = AutoMergingRetriever(document_store=nonclosable_document_store) + retriever.close() + assert nonclosable_document_store.mock_calls == [] diff --git a/test/components/retrievers/test_auto_merging_retriever_async.py b/test/components/retrievers/test_auto_merging_retriever_async.py index 21e1741db52..71799d702a6 100644 --- a/test/components/retrievers/test_auto_merging_retriever_async.py +++ b/test/components/retrievers/test_auto_merging_retriever_async.py @@ -2,6 +2,8 @@ # # SPDX-License-Identifier: Apache-2.0 +from unittest.mock import AsyncMock, Mock + import pytest from haystack import Document @@ -206,3 +208,16 @@ async def test_run_go_up_hierarchy_multiple_levels_hit_root_document(self, in_me assert len(result["documents"]) == 1 assert result["documents"][0].meta["__level"] == 0 # hit root document + + @pytest.mark.asyncio + async def test_close_async(self): + closable_document_store = Mock(spec=["close_async"]) + closable_document_store.close_async = AsyncMock() + retriever = AutoMergingRetriever(document_store=closable_document_store) + await retriever.close_async() + closable_document_store.close_async.assert_awaited_once_with() + + nonclosable_document_store = Mock(spec=[]) + retriever = AutoMergingRetriever(document_store=nonclosable_document_store) + await retriever.close_async() + assert nonclosable_document_store.mock_calls == [] diff --git a/test/components/retrievers/test_filter_retriever.py b/test/components/retrievers/test_filter_retriever.py index 33f15bcfc14..6974b4ed048 100644 --- a/test/components/retrievers/test_filter_retriever.py +++ b/test/components/retrievers/test_filter_retriever.py @@ -3,6 +3,7 @@ # SPDX-License-Identifier: Apache-2.0 from typing import Any +from unittest.mock import Mock import pytest @@ -144,3 +145,14 @@ def test_run_with_pipeline(self, sample_document_store, sample_docs): results_docs = result["retriever"]["documents"] assert results_docs assert TestFilterRetriever._documents_equal(results_docs, sample_docs["en_docs"]) + + def test_close(self): + closable_document_store = Mock(spec=["close"]) + retriever = FilterRetriever(document_store=closable_document_store) + retriever.close() + closable_document_store.close.assert_called_once_with() + + nonclosable_document_store = Mock(spec=[]) + retriever = FilterRetriever(document_store=nonclosable_document_store) + retriever.close() + assert nonclosable_document_store.mock_calls == [] diff --git a/test/components/retrievers/test_filter_retriever_async.py b/test/components/retrievers/test_filter_retriever_async.py index db11addbf59..1786cbbcf48 100644 --- a/test/components/retrievers/test_filter_retriever_async.py +++ b/test/components/retrievers/test_filter_retriever_async.py @@ -3,6 +3,7 @@ # SPDX-License-Identifier: Apache-2.0 from typing import Any +from unittest.mock import AsyncMock, Mock import pytest @@ -94,3 +95,16 @@ async def test_run_with_pipeline(self, sample_document_store, sample_docs): results_docs = result["retriever"]["documents"] assert results_docs assert TestFilterRetrieverAsync._documents_equal(results_docs, sample_docs["en_docs"]) + + @pytest.mark.asyncio + async def test_close_async(self): + closable_document_store = Mock(spec=["close_async"]) + closable_document_store.close_async = AsyncMock() + retriever = FilterRetriever(document_store=closable_document_store) + await retriever.close_async() + closable_document_store.close_async.assert_awaited_once_with() + + nonclosable_document_store = Mock(spec=[]) + retriever = FilterRetriever(document_store=nonclosable_document_store) + await retriever.close_async() + assert nonclosable_document_store.mock_calls == [] diff --git a/test/components/retrievers/test_sentence_window_retriever.py b/test/components/retrievers/test_sentence_window_retriever.py index 07cd8f7dbd3..d7cb2227290 100644 --- a/test/components/retrievers/test_sentence_window_retriever.py +++ b/test/components/retrievers/test_sentence_window_retriever.py @@ -4,7 +4,7 @@ import random import re -from unittest.mock import ANY +from unittest.mock import ANY, Mock import pytest @@ -329,3 +329,14 @@ def test_serialization_deserialization_in_pipeline(self, in_memory_doc_store): deserialized = Pipeline.from_dict(serialized) assert deserialized == pipe + + def test_close(self): + closable_document_store = Mock(spec=["close"]) + retriever = SentenceWindowRetriever(document_store=closable_document_store) + retriever.close() + closable_document_store.close.assert_called_once_with() + + nonclosable_document_store = Mock(spec=[]) + retriever = SentenceWindowRetriever(document_store=nonclosable_document_store) + retriever.close() + assert nonclosable_document_store.mock_calls == [] diff --git a/test/components/retrievers/test_sentence_window_retriever_async.py b/test/components/retrievers/test_sentence_window_retriever_async.py index d1c8306b4a8..3c776f5b9f3 100644 --- a/test/components/retrievers/test_sentence_window_retriever_async.py +++ b/test/components/retrievers/test_sentence_window_retriever_async.py @@ -4,6 +4,7 @@ import random import re +from unittest.mock import AsyncMock, Mock import pytest @@ -224,3 +225,16 @@ async def test_serialization_deserialization_in_pipeline(self, in_memory_doc_sto deserialized = Pipeline.from_dict(serialized) assert deserialized == pipe + + @pytest.mark.asyncio + async def test_close_async(self): + closable_document_store = Mock(spec=["close_async"]) + closable_document_store.close_async = AsyncMock() + retriever = SentenceWindowRetriever(document_store=closable_document_store) + await retriever.close_async() + closable_document_store.close_async.assert_awaited_once_with() + + nonclosable_document_store = Mock(spec=[]) + retriever = SentenceWindowRetriever(document_store=nonclosable_document_store) + await retriever.close_async() + assert nonclosable_document_store.mock_calls == [] diff --git a/test/components/writers/test_document_writer.py b/test/components/writers/test_document_writer.py index 9c12c9d9923..36fd2430274 100644 --- a/test/components/writers/test_document_writer.py +++ b/test/components/writers/test_document_writer.py @@ -2,6 +2,8 @@ # # SPDX-License-Identifier: Apache-2.0 +from unittest.mock import AsyncMock, Mock + import pytest from haystack import Document @@ -107,6 +109,17 @@ def test_run_skip_policy(self, in_memory_doc_store): result = writer.run(documents=documents) assert result["documents_written"] == 0 + def test_close(self): + closable_document_store = Mock(spec=["close"]) + writer = DocumentWriter(document_store=closable_document_store) + writer.close() + closable_document_store.close.assert_called_once_with() + + nonclosable_document_store = Mock(spec=[]) + writer = DocumentWriter(document_store=nonclosable_document_store) + writer.close() + assert nonclosable_document_store.mock_calls == [] + @pytest.mark.asyncio async def test_run_async_invalid_docstore(self): mocked_docstore_class = document_store_class("MockedDocumentStore") @@ -144,3 +157,16 @@ async def test_run_async_skip_policy(self, in_memory_doc_store): result = await writer.run_async(documents=documents) assert result["documents_written"] == 0 + + @pytest.mark.asyncio + async def test_close_async(self): + closable_document_store = Mock(spec=["close_async"]) + closable_document_store.close_async = AsyncMock() + writer = DocumentWriter(document_store=closable_document_store) + await writer.close_async() + closable_document_store.close_async.assert_awaited_once_with() + + nonclosable_document_store = Mock(spec=[]) + writer = DocumentWriter(document_store=nonclosable_document_store) + await writer.close_async() + assert nonclosable_document_store.mock_calls == []