Skip to content

Commit 076bc83

Browse files
authored
feat: Azure AI Search - add closing methods (#3662)
1 parent 8f3fabf commit 076bc83

8 files changed

Lines changed: 110 additions & 0 deletions

File tree

integrations/azure_ai_search/src/haystack_integrations/components/retrievers/azure_ai_search/bm25_retriever.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -96,6 +96,12 @@ def from_dict(cls, data: dict[str, Any]) -> "AzureAISearchBM25Retriever":
9696
data["init_parameters"]["filter_policy"] = FilterPolicy.from_str(data["init_parameters"]["filter_policy"])
9797
return default_from_dict(cls, data)
9898

99+
def close(self) -> None:
100+
"""
101+
Release the synchronous resources of the underlying Document Store.
102+
"""
103+
self._document_store.close()
104+
99105
@component.output_types(documents=list[Document])
100106
def run(
101107
self, query: str, filters: dict[str, Any] | None = None, top_k: int | None = None

integrations/azure_ai_search/src/haystack_integrations/components/retrievers/azure_ai_search/embedding_retriever.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,6 +93,12 @@ def from_dict(cls, data: dict[str, Any]) -> "AzureAISearchEmbeddingRetriever":
9393
data["init_parameters"]["filter_policy"] = FilterPolicy.from_str(data["init_parameters"]["filter_policy"])
9494
return default_from_dict(cls, data)
9595

96+
def close(self) -> None:
97+
"""
98+
Release the synchronous resources of the underlying Document Store.
99+
"""
100+
self._document_store.close()
101+
96102
@component.output_types(documents=list[Document])
97103
def run(
98104
self, query_embedding: list[float], filters: dict[str, Any] | None = None, top_k: int | None = None

integrations/azure_ai_search/src/haystack_integrations/components/retrievers/azure_ai_search/hybrid_retriever.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -96,6 +96,12 @@ def from_dict(cls, data: dict[str, Any]) -> "AzureAISearchHybridRetriever":
9696
data["init_parameters"]["filter_policy"] = FilterPolicy.from_str(data["init_parameters"]["filter_policy"])
9797
return default_from_dict(cls, data)
9898

99+
def close(self) -> None:
100+
"""
101+
Release the synchronous resources of the underlying Document Store.
102+
"""
103+
self._document_store.close()
104+
99105
@component.output_types(documents=list[Document])
100106
def run(
101107
self,

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

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44

55
import logging as python_logging
66
from collections.abc import Mapping, Sequence
7+
from contextlib import suppress
78
from datetime import datetime
89
from typing import Any
910

@@ -387,6 +388,19 @@ def from_dict(cls, data: dict[str, Any]) -> "AzureAISearchDocumentStore":
387388
data["init_parameters"]["vector_search_configuration"] = VectorSearch(vector_search_configuration)
388389
return default_from_dict(cls, data)
389390

391+
def close(self) -> None:
392+
"""
393+
Release the associated synchronous resources.
394+
"""
395+
if self._client is not None:
396+
with suppress(Exception):
397+
self._client.close()
398+
self._client = None
399+
if self._index_client is not None:
400+
with suppress(Exception):
401+
self._index_client.close()
402+
self._index_client = None
403+
390404
def count_documents(self) -> int:
391405
"""
392406
Returns how many documents are present in the search index.

integrations/azure_ai_search/tests/test_bm25_retriever.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -97,6 +97,16 @@ def test_from_dict():
9797
assert retriever._filter_policy == FilterPolicy.REPLACE
9898

9999

100+
def test_close():
101+
mock_store = Mock(spec=AzureAISearchDocumentStore)
102+
retriever = AzureAISearchBM25Retriever(document_store=mock_store)
103+
104+
retriever.close()
105+
106+
mock_store.close.assert_called_once()
107+
assert retriever._document_store is mock_store
108+
109+
100110
def test_run():
101111
mock_store = Mock(spec=AzureAISearchDocumentStore)
102112
mock_store._bm25_retrieval.return_value = [Document(content="Test doc")]

integrations/azure_ai_search/tests/test_document_store.py

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -261,6 +261,47 @@ def test_init():
261261
assert document_store._vector_search_configuration == DEFAULT_VECTOR_SEARCH
262262

263263

264+
def test_close():
265+
store = AzureAISearchDocumentStore(
266+
api_key=Secret.from_token("fake-api-key"),
267+
azure_endpoint=Secret.from_token("fake-endpoint"),
268+
)
269+
client = Mock()
270+
index_client = Mock()
271+
store._client = client
272+
store._index_client = index_client
273+
274+
store.close()
275+
276+
client.close.assert_called_once()
277+
index_client.close.assert_called_once()
278+
assert store._client is None
279+
assert store._index_client is None
280+
281+
store.close()
282+
283+
client.close.assert_called_once()
284+
index_client.close.assert_called_once()
285+
286+
287+
def test_close_is_exception_safe():
288+
store = AzureAISearchDocumentStore(
289+
api_key=Secret.from_token("fake-api-key"),
290+
azure_endpoint=Secret.from_token("fake-endpoint"),
291+
)
292+
client = Mock()
293+
client.close.side_effect = RuntimeError("boom")
294+
index_client = Mock()
295+
index_client.close.side_effect = RuntimeError("boom")
296+
store._client = client
297+
store._index_client = index_client
298+
299+
store.close()
300+
301+
assert store._client is None
302+
assert store._index_client is None
303+
304+
264305
def test_token_credential_takes_priority_over_api_key(monkeypatch: pytest.MonkeyPatch) -> None:
265306
monkeypatch.setenv("AZURE_AI_SEARCH_API_KEY", "test-api-key")
266307
monkeypatch.setenv("AZURE_AI_SEARCH_ENDPOINT", "test-endpoint")
@@ -496,6 +537,13 @@ class TestDocumentStore(
496537
def assert_documents_are_equal(self, received: list[Document], expected: list[Document]):
497538
_assert_documents_are_equal(received, expected)
498539

540+
def test_close_and_reopen(self, document_store: AzureAISearchDocumentStore):
541+
assert document_store.count_documents() == 0
542+
document_store.close()
543+
assert document_store._client is None
544+
assert document_store._index_client is None
545+
assert document_store.count_documents() == 0
546+
499547
def test_write_documents(self, document_store: AzureAISearchDocumentStore):
500548
docs = [Document(id="1")]
501549
assert document_store.write_documents(docs) == 1

integrations/azure_ai_search/tests/test_embedding_retriever.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,16 @@ def test_from_dict():
100100
assert retriever._filter_policy == FilterPolicy.REPLACE
101101

102102

103+
def test_close():
104+
mock_store = Mock(spec=AzureAISearchDocumentStore)
105+
retriever = AzureAISearchEmbeddingRetriever(document_store=mock_store)
106+
107+
retriever.close()
108+
109+
mock_store.close.assert_called_once()
110+
assert retriever._document_store is mock_store
111+
112+
103113
def test_run():
104114
mock_store = Mock(spec=AzureAISearchDocumentStore)
105115
mock_store._embedding_retrieval.return_value = [Document(content="Test doc", embedding=[0.1, 0.2])]

integrations/azure_ai_search/tests/test_hybrid_retriever.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,16 @@ def test_from_dict():
100100
assert retriever._filter_policy == FilterPolicy.REPLACE
101101

102102

103+
def test_close():
104+
mock_store = Mock(spec=AzureAISearchDocumentStore)
105+
retriever = AzureAISearchHybridRetriever(document_store=mock_store)
106+
107+
retriever.close()
108+
109+
mock_store.close.assert_called_once()
110+
assert retriever._document_store is mock_store
111+
112+
103113
def test_run():
104114
mock_store = Mock(spec=AzureAISearchDocumentStore)
105115
mock_store._hybrid_retrieval.return_value = [Document(content="Test doc", embedding=[0.1, 0.2])]

0 commit comments

Comments
 (0)