Skip to content

Commit 0a26a74

Browse files
authored
feat: arcadedb - add closing methods (#3653)
1 parent fa4cedf commit 0a26a74

4 files changed

Lines changed: 111 additions & 4 deletions

File tree

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

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -122,6 +122,12 @@ def to_dict(self) -> dict[str, Any]:
122122
filter_policy=self._filter_policy.value,
123123
)
124124

125+
def close(self) -> None:
126+
"""
127+
Release the synchronous resources of the underlying Document Store.
128+
"""
129+
self._document_store.close()
130+
125131
@classmethod
126132
def from_dict(cls, data: dict[str, Any]) -> "ArcadeDBEmbeddingRetriever":
127133
"""

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

Lines changed: 19 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
"""ArcadeDB DocumentStore for Haystack 2.x — document storage + vector search via HTTP/JSON API."""
66

77
import logging
8+
from contextlib import suppress
89
from http import HTTPStatus
910
from typing import Any, ClassVar
1011

@@ -95,7 +96,7 @@ def __init__(
9596
self._recreate_type = recreate_type
9697
self._create_database = create_database
9798

98-
self._session = requests.Session()
99+
self._session: requests.Session | None = None
99100
self._initialized = False
100101

101102
def to_dict(self) -> dict[str, Any]:
@@ -134,10 +135,25 @@ def from_dict(cls, data: dict[str, Any]) -> "ArcadeDBDocumentStore":
134135
init_params[key] = Secret.from_dict(init_params[key])
135136
return default_from_dict(cls, data)
136137

138+
def close(self) -> None:
139+
"""
140+
Release the associated synchronous resources.
141+
"""
142+
if self._session is not None:
143+
with suppress(Exception):
144+
self._session.close()
145+
self._session = None
146+
137147
# ------------------------------------------------------------------
138148
# HTTP helpers
139149
# ------------------------------------------------------------------
140150

151+
def _ensure_session(self) -> requests.Session:
152+
"""Lazily create the HTTP session"""
153+
if self._session is None:
154+
self._session = requests.Session()
155+
return self._session
156+
141157
def _auth(self) -> tuple[str, str] | None:
142158
user = self._username.resolve_value() if self._username else None
143159
pwd = self._password.resolve_value() if self._password else None
@@ -152,7 +168,7 @@ def _command(self, sql: str, *, positional_params: list[Any] | None = None) -> l
152168
if positional_params:
153169
payload["params"] = positional_params
154170

155-
resp = self._session.post(url, json=payload, auth=self._auth())
171+
resp = self._ensure_session().post(url, json=payload, auth=self._auth())
156172
if resp.status_code >= HTTPStatus.BAD_REQUEST:
157173
msg = f"ArcadeDB command failed ({resp.status_code}): {resp.text}"
158174
raise RuntimeError(msg)
@@ -163,7 +179,7 @@ def _command(self, sql: str, *, positional_params: list[Any] | None = None) -> l
163179
def _server_command(self, command: str) -> dict[str, Any]:
164180
"""Execute a server-level command (e.g. CREATE DATABASE)."""
165181
url = f"{self._url}/api/v1/server"
166-
resp = self._session.post(url, json={"command": command}, auth=self._auth())
182+
resp = self._ensure_session().post(url, json={"command": command}, auth=self._auth())
167183
if resp.status_code >= HTTPStatus.BAD_REQUEST:
168184
msg = f"ArcadeDB server command failed ({resp.status_code}): {resp.text}"
169185
raise RuntimeError(msg)

integrations/arcadedb/tests/test_document_store.py

Lines changed: 24 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
import dataclasses
66
import datetime
77
import os
8-
from unittest.mock import MagicMock
8+
from unittest.mock import MagicMock, Mock
99

1010
import pytest
1111
from haystack import Document
@@ -338,6 +338,20 @@ def test_embedding_retrieval_without_filter_returns_all(self, store):
338338
docs = store._embedding_retrieval([0.0] * 4)
339339
assert [d.id for d in docs] == ["a"]
340340

341+
def test_close(self, store):
342+
store._session = Mock()
343+
session = store._session
344+
store.close()
345+
session.close.assert_called_once_with()
346+
assert store._session is None
347+
store.close()
348+
349+
def test_close_is_exception_safe(self, store):
350+
store._session = Mock()
351+
store._session.close.side_effect = RuntimeError
352+
store.close()
353+
assert store._session is None
354+
341355

342356
@pytest.mark.skipif(
343357
not os.environ.get("ARCADEDB_PASSWORD"),
@@ -388,6 +402,15 @@ def assert_documents_are_equal(self, received: list[Document], expected: list[Do
388402
expected_clean = dataclasses.replace(expected_doc, embedding=None)
389403
assert actual == expected_clean
390404

405+
def test_close_and_reopen(self, document_store: ArcadeDBDocumentStore):
406+
document_store.write_documents([Document(id="1")])
407+
assert document_store.count_documents() == 1
408+
409+
document_store.close()
410+
assert document_store._session is None
411+
412+
assert document_store.count_documents() == 1
413+
391414
def test_write_documents(self, document_store: ArcadeDBDocumentStore):
392415
"""Override mixin: test default write_documents and duplicate fail behaviour."""
393416
docs = [Document(id="1")]
Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
1+
# SPDX-FileCopyrightText: 2025-present deepset GmbH <info@deepset.ai>
2+
#
3+
# SPDX-License-Identifier: Apache-2.0
4+
5+
from unittest.mock import Mock
6+
7+
from haystack.dataclasses import Document
8+
from haystack.document_stores.types import FilterPolicy
9+
10+
from haystack_integrations.components.retrievers.arcadedb import ArcadeDBEmbeddingRetriever
11+
from haystack_integrations.document_stores.arcadedb import ArcadeDBDocumentStore
12+
13+
14+
class TestEmbeddingRetriever:
15+
def test_init_default(self):
16+
mock_store = Mock(spec=ArcadeDBDocumentStore)
17+
retriever = ArcadeDBEmbeddingRetriever(document_store=mock_store)
18+
assert retriever._document_store == mock_store
19+
assert retriever._filters is None
20+
assert retriever._top_k == 10
21+
assert retriever._filter_policy == FilterPolicy.REPLACE
22+
23+
def test_init(self):
24+
mock_store = Mock(spec=ArcadeDBDocumentStore)
25+
retriever = ArcadeDBEmbeddingRetriever(
26+
document_store=mock_store, filters={"field": "value"}, top_k=5, filter_policy=FilterPolicy.MERGE
27+
)
28+
assert retriever._filters == {"field": "value"}
29+
assert retriever._top_k == 5
30+
assert retriever._filter_policy == FilterPolicy.MERGE
31+
32+
def test_to_dict_from_dict(self):
33+
store = ArcadeDBDocumentStore(url="http://localhost:2480", database="test", create_database=False)
34+
retriever = ArcadeDBEmbeddingRetriever(document_store=store, filters={"field": "value"}, top_k=5)
35+
36+
restored = ArcadeDBEmbeddingRetriever.from_dict(retriever.to_dict())
37+
38+
assert isinstance(restored._document_store, ArcadeDBDocumentStore)
39+
assert restored._document_store._database == "test"
40+
assert restored._filters == {"field": "value"}
41+
assert restored._top_k == 5
42+
assert restored._filter_policy == FilterPolicy.REPLACE
43+
44+
def test_close(self):
45+
mock_store = Mock(spec=ArcadeDBDocumentStore)
46+
retriever = ArcadeDBEmbeddingRetriever(document_store=mock_store)
47+
48+
retriever.close()
49+
50+
mock_store.close.assert_called_once()
51+
assert retriever._document_store is mock_store
52+
53+
def test_run(self):
54+
mock_store = Mock(spec=ArcadeDBDocumentStore)
55+
doc = Document(content="Test doc", embedding=[0.1, 0.2])
56+
mock_store._embedding_retrieval.return_value = [doc]
57+
58+
retriever = ArcadeDBEmbeddingRetriever(document_store=mock_store)
59+
res = retriever.run(query_embedding=[0.3, 0.5])
60+
61+
mock_store._embedding_retrieval.assert_called_once_with(query_embedding=[0.3, 0.5], filters=None, top_k=10)
62+
assert res == {"documents": [doc]}

0 commit comments

Comments
 (0)