weaviate-haystack 2.0.0__tar.gz → 2.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (22) hide show
  1. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/PKG-INFO +2 -1
  2. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/pydoc/config.yml +1 -1
  3. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/pyproject.toml +1 -0
  4. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/src/haystack_integrations/document_stores/weaviate/document_store.py +66 -41
  5. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/tests/test_document_store.py +28 -4
  6. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/.gitignore +0 -0
  7. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/LICENSE.txt +0 -0
  8. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/README.md +0 -0
  9. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/docker-compose.yml +0 -0
  10. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/src/haystack_integrations/components/retrievers/weaviate/__init__.py +0 -0
  11. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/src/haystack_integrations/components/retrievers/weaviate/bm25_retriever.py +0 -0
  12. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/src/haystack_integrations/components/retrievers/weaviate/embedding_retriever.py +0 -0
  13. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/src/haystack_integrations/document_stores/weaviate/__init__.py +0 -0
  14. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/src/haystack_integrations/document_stores/weaviate/_filters.py +0 -0
  15. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/src/haystack_integrations/document_stores/weaviate/auth.py +0 -0
  16. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/tests/__init__.py +0 -0
  17. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/tests/conftest.py +0 -0
  18. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/tests/test_auth.py +0 -0
  19. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/tests/test_bm25_retriever.py +0 -0
  20. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/tests/test_embedding_retriever.py +0 -0
  21. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/tests/test_files/robot1.jpg +0 -0
  22. {weaviate_haystack-2.0.0 → weaviate_haystack-2.1.0}/tests/test_filters.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: weaviate-haystack
3
- Version: 2.0.0
3
+ Version: 2.1.0
4
4
  Summary: An integration of Weaviate vector database with Haystack
5
5
  Project-URL: Source, https://github.com/deepset-ai/haystack-core-integrations
6
6
  Project-URL: Documentation, https://github.com/deepset-ai/haystack-core-integrations/blob/main/integrations/weaviate/README.md
@@ -9,6 +9,7 @@ Author-email: deepset GmbH <info@deepset.ai>
9
9
  License-Expression: Apache-2.0
10
10
  License-File: LICENSE.txt
11
11
  Classifier: Development Status :: 4 - Beta
12
+ Classifier: License :: OSI Approved :: Apache Software License
12
13
  Classifier: Programming Language :: Python
13
14
  Classifier: Programming Language :: Python :: 3.8
14
15
  Classifier: Programming Language :: Python :: 3.9
@@ -18,7 +18,7 @@ processors:
18
18
  - type: smart
19
19
  - type: crossref
20
20
  renderer:
21
- type: haystack_pydoc_tools.renderers.ReadmePreviewRenderer
21
+ type: haystack_pydoc_tools.renderers.ReadmeIntegrationRenderer
22
22
  excerpt: Weaviate integration for Haystack
23
23
  category_slug: integrations-api
24
24
  title: Weaviate
@@ -12,6 +12,7 @@ license = "Apache-2.0"
12
12
  keywords = []
13
13
  authors = [{ name = "deepset GmbH", email = "info@deepset.ai" }]
14
14
  classifiers = [
15
+ "License :: OSI Approved :: Apache Software License",
15
16
  "Development Status :: 4 - Beta",
16
17
  "Programming Language :: Python",
17
18
  "Programming Language :: Python :: 3.8",
@@ -139,47 +139,72 @@ class WeaviateDocumentStore:
139
139
  :param grpc_secure:
140
140
  Whether to use a secure channel for the underlying gRPC API.
141
141
  """
142
+ self._url = url
143
+ self._auth_client_secret = auth_client_secret
144
+ self._additional_headers = additional_headers
145
+ self._embedded_options = embedded_options
146
+ self._additional_config = additional_config
147
+ self._grpc_port = grpc_port
148
+ self._grpc_secure = grpc_secure
149
+ self._client = None
150
+ self._collection = None
151
+ # Store the connection settings dictionary
152
+ self._collection_settings = collection_settings or {
153
+ "class": "Default",
154
+ "invertedIndexConfig": {"indexNullState": True},
155
+ "properties": DOCUMENT_COLLECTION_PROPERTIES,
156
+ }
157
+ self._clean_connection_settings()
158
+
159
+ def _clean_connection_settings(self):
160
+ # Set the class if not set
161
+ _class_name = self._collection_settings.get("class", "Default")
162
+ _class_name = _class_name[0].upper() + _class_name[1:]
163
+ self._collection_settings["class"] = _class_name
164
+ # Set the properties if they're not set
165
+ self._collection_settings["properties"] = self._collection_settings.get(
166
+ "properties", DOCUMENT_COLLECTION_PROPERTIES
167
+ )
168
+
169
+ @property
170
+ def client(self):
171
+ if self._client:
172
+ return self._client
173
+
142
174
  # proxies, timeout_config, trust_env are part of additional_config now
143
175
  # startup_period has been removed
144
176
  self._client = weaviate.WeaviateClient(
145
177
  connection_params=(
146
- weaviate.connect.base.ConnectionParams.from_url(url=url, grpc_port=grpc_port, grpc_secure=grpc_secure)
147
- if url
178
+ weaviate.connect.base.ConnectionParams.from_url(
179
+ url=self._url, grpc_port=self._grpc_port, grpc_secure=self._grpc_secure
180
+ )
181
+ if self._url
148
182
  else None
149
183
  ),
150
- auth_client_secret=auth_client_secret.resolve_value() if auth_client_secret else None,
151
- additional_config=additional_config,
152
- additional_headers=additional_headers,
153
- embedded_options=embedded_options,
184
+ auth_client_secret=self._auth_client_secret.resolve_value() if self._auth_client_secret else None,
185
+ additional_config=self._additional_config,
186
+ additional_headers=self._additional_headers,
187
+ embedded_options=self._embedded_options,
154
188
  skip_init_checks=False,
155
189
  )
190
+
156
191
  self._client.connect()
157
192
 
158
193
  # Test connection, it will raise an exception if it fails.
159
194
  self._client.collections._get_all(simple=True)
195
+ if not self._client.collections.exists(self._collection_settings["class"]):
196
+ self._client.collections.create_from_dict(self._collection_settings)
160
197
 
161
- if collection_settings is None:
162
- collection_settings = {
163
- "class": "Default",
164
- "invertedIndexConfig": {"indexNullState": True},
165
- "properties": DOCUMENT_COLLECTION_PROPERTIES,
166
- }
167
- else:
168
- # Set the class if not set
169
- collection_settings["class"] = collection_settings.get("class", "default").capitalize()
170
- # Set the properties if they're not set
171
- collection_settings["properties"] = collection_settings.get("properties", DOCUMENT_COLLECTION_PROPERTIES)
198
+ return self._client
172
199
 
173
- if not self._client.collections.exists(collection_settings["class"]):
174
- self._client.collections.create_from_dict(collection_settings)
200
+ @property
201
+ def collection(self):
202
+ if self._collection:
203
+ return self._collection
175
204
 
176
- self._url = url
177
- self._collection_settings = collection_settings
178
- self._auth_client_secret = auth_client_secret
179
- self._additional_headers = additional_headers
180
- self._embedded_options = embedded_options
181
- self._additional_config = additional_config
182
- self._collection = self._client.collections.get(collection_settings["class"])
205
+ client = self.client
206
+ self._collection = client.collections.get(self._collection_settings["class"])
207
+ return self._collection
183
208
 
184
209
  def to_dict(self) -> Dict[str, Any]:
185
210
  """
@@ -228,7 +253,7 @@ class WeaviateDocumentStore:
228
253
  """
229
254
  Returns the number of documents present in the DocumentStore.
230
255
  """
231
- total = self._collection.aggregate.over_all(total_count=True).total_count
256
+ total = self.collection.aggregate.over_all(total_count=True).total_count
232
257
  return total if total else 0
233
258
 
234
259
  def _to_data_object(self, document: Document) -> Dict[str, Any]:
@@ -300,16 +325,16 @@ class WeaviateDocumentStore:
300
325
  return Document.from_dict(document_data)
301
326
 
302
327
  def _query(self) -> List[Dict[str, Any]]:
303
- properties = [p.name for p in self._collection.config.get().properties]
328
+ properties = [p.name for p in self.collection.config.get().properties]
304
329
  try:
305
- result = self._collection.iterator(include_vector=True, return_properties=properties)
330
+ result = self.collection.iterator(include_vector=True, return_properties=properties)
306
331
  except weaviate.exceptions.WeaviateQueryError as e:
307
332
  msg = f"Failed to query documents in Weaviate. Error: {e.message}"
308
333
  raise DocumentStoreError(msg) from e
309
334
  return result
310
335
 
311
336
  def _query_with_filters(self, filters: Dict[str, Any]) -> List[Dict[str, Any]]:
312
- properties = [p.name for p in self._collection.config.get().properties]
337
+ properties = [p.name for p in self.collection.config.get().properties]
313
338
  # When querying with filters we need to paginate using limit and offset as using
314
339
  # a cursor with after is not possible. See the official docs:
315
340
  # https://weaviate.io/developers/weaviate/api/graphql/additional-operators#cursor-with-after
@@ -325,7 +350,7 @@ class WeaviateDocumentStore:
325
350
  # Keep querying until we get all documents matching the filters
326
351
  while partial_result is None or len(partial_result.objects) == DEFAULT_QUERY_LIMIT:
327
352
  try:
328
- partial_result = self._collection.query.fetch_objects(
353
+ partial_result = self.collection.query.fetch_objects(
329
354
  filters=convert_filters(filters),
330
355
  include_vector=True,
331
356
  limit=DEFAULT_QUERY_LIMIT,
@@ -363,7 +388,7 @@ class WeaviateDocumentStore:
363
388
  Raises in case of errors.
364
389
  """
365
390
 
366
- with self._client.batch.dynamic() as batch:
391
+ with self.client.batch.dynamic() as batch:
367
392
  for doc in documents:
368
393
  if not isinstance(doc, Document):
369
394
  msg = f"Expected a Document, got '{type(doc)}' instead."
@@ -371,11 +396,11 @@ class WeaviateDocumentStore:
371
396
 
372
397
  batch.add_object(
373
398
  properties=self._to_data_object(doc),
374
- collection=self._collection.name,
399
+ collection=self.collection.name,
375
400
  uuid=generate_uuid5(doc.id),
376
401
  vector=doc.embedding,
377
402
  )
378
- if failed_objects := self._client.batch.failed_objects:
403
+ if failed_objects := self.client.batch.failed_objects:
379
404
  # We fallback to use the UUID if the _original_id is not present, this is just to be
380
405
  mapped_objects = {}
381
406
  for obj in failed_objects:
@@ -411,12 +436,12 @@ class WeaviateDocumentStore:
411
436
  msg = f"Expected a Document, got '{type(doc)}' instead."
412
437
  raise ValueError(msg)
413
438
 
414
- if policy == DuplicatePolicy.SKIP and self._collection.data.exists(uuid=generate_uuid5(doc.id)):
439
+ if policy == DuplicatePolicy.SKIP and self.collection.data.exists(uuid=generate_uuid5(doc.id)):
415
440
  # This Document already exists, we skip it
416
441
  continue
417
442
 
418
443
  try:
419
- self._collection.data.insert(
444
+ self.collection.data.insert(
420
445
  uuid=generate_uuid5(doc.id),
421
446
  properties=self._to_data_object(doc),
422
447
  vector=doc.embedding,
@@ -452,13 +477,13 @@ class WeaviateDocumentStore:
452
477
  :param document_ids: The object_ids to delete.
453
478
  """
454
479
  weaviate_ids = [generate_uuid5(doc_id) for doc_id in document_ids]
455
- self._collection.data.delete_many(where=weaviate.classes.query.Filter.by_id().contains_any(weaviate_ids))
480
+ self.collection.data.delete_many(where=weaviate.classes.query.Filter.by_id().contains_any(weaviate_ids))
456
481
 
457
482
  def _bm25_retrieval(
458
483
  self, query: str, filters: Optional[Dict[str, Any]] = None, top_k: Optional[int] = None
459
484
  ) -> List[Document]:
460
- properties = [p.name for p in self._collection.config.get().properties]
461
- result = self._collection.query.bm25(
485
+ properties = [p.name for p in self.collection.config.get().properties]
486
+ result = self.collection.query.bm25(
462
487
  query=query,
463
488
  filters=convert_filters(filters) if filters else None,
464
489
  limit=top_k,
@@ -482,8 +507,8 @@ class WeaviateDocumentStore:
482
507
  msg = "Can't use 'distance' and 'certainty' parameters together"
483
508
  raise ValueError(msg)
484
509
 
485
- properties = [p.name for p in self._collection.config.get().properties]
486
- result = self._collection.query.near_vector(
510
+ properties = [p.name for p in self.collection.config.get().properties]
511
+ result = self.collection.query.near_vector(
487
512
  near_vector=query_embedding,
488
513
  distance=distance,
489
514
  certainty=certainty,
@@ -38,6 +38,12 @@ from weaviate.embedded import (
38
38
  )
39
39
 
40
40
 
41
+ @patch("haystack_integrations.document_stores.weaviate.document_store.weaviate.WeaviateClient")
42
+ def test_init_is_lazy(_mock_client):
43
+ _ = WeaviateDocumentStore()
44
+ _mock_client.assert_not_called()
45
+
46
+
41
47
  @pytest.mark.integration
42
48
  class TestWeaviateDocumentStore(CountDocumentsTest, WriteDocumentsTest, DeleteDocumentsTest, FilterDocumentsTest):
43
49
  @pytest.fixture
@@ -57,7 +63,7 @@ class TestWeaviateDocumentStore(CountDocumentsTest, WriteDocumentsTest, DeleteDo
57
63
  collection_settings=collection_settings,
58
64
  )
59
65
  yield store
60
- store._client.collections.delete(collection_settings["class"])
66
+ store.client.collections.delete(collection_settings["class"])
61
67
 
62
68
  @pytest.fixture
63
69
  def filterable_docs(self) -> List[Document]:
@@ -150,12 +156,12 @@ class TestWeaviateDocumentStore(CountDocumentsTest, WriteDocumentsTest, DeleteDo
150
156
  assert received_meta.get(key) == expected_meta.get(key)
151
157
 
152
158
  @patch("haystack_integrations.document_stores.weaviate.document_store.weaviate.WeaviateClient")
153
- def test_init(self, mock_weaviate_client_class, monkeypatch):
159
+ def test_connection(self, mock_weaviate_client_class, monkeypatch):
154
160
  mock_client = MagicMock()
155
161
  mock_client.collections.exists.return_value = False
156
162
  mock_weaviate_client_class.return_value = mock_client
157
163
  monkeypatch.setenv("WEAVIATE_API_KEY", "my_api_key")
158
- WeaviateDocumentStore(
164
+ ds = WeaviateDocumentStore(
159
165
  collection_settings={"class": "My_collection"},
160
166
  auth_client_secret=AuthApiKey(),
161
167
  additional_headers={"X-HuggingFace-Api-Key": "MY_HUGGINGFACE_KEY"},
@@ -170,8 +176,11 @@ class TestWeaviateDocumentStore(CountDocumentsTest, WriteDocumentsTest, DeleteDo
170
176
  ),
171
177
  )
172
178
 
173
- # Verify client is created with correct parameters
179
+ # Trigger the actual database connection by accessing the `client` property so we
180
+ # can assert the setup was good
181
+ _ = ds.client
174
182
 
183
+ # Verify client is created with correct parameters
175
184
  mock_weaviate_client_class.assert_called_once_with(
176
185
  auth_client_secret=AuthApiKey().resolve_value(),
177
186
  connection_params=None,
@@ -664,3 +673,18 @@ class TestWeaviateDocumentStore(CountDocumentsTest, WriteDocumentsTest, DeleteDo
664
673
  document_store.write_documents(docs)
665
674
  with pytest.raises(DocumentStoreError):
666
675
  document_store.filter_documents({"field": "content", "operator": "==", "value": "This is some content"})
676
+
677
+ def test_schema_class_name_conversion_preserves_pascal_case(self):
678
+ collection_settings = {"class": "CaseDocument"}
679
+ doc_score = WeaviateDocumentStore(
680
+ url="http://localhost:8080",
681
+ collection_settings=collection_settings,
682
+ )
683
+ assert doc_score._collection_settings["class"] == "CaseDocument"
684
+
685
+ collection_settings = {"class": "lower_case_name"}
686
+ doc_score = WeaviateDocumentStore(
687
+ url="http://localhost:8080",
688
+ collection_settings=collection_settings,
689
+ )
690
+ assert doc_score._collection_settings["class"] == "Lower_case_name"