codekingpro/portable-devtools
114k
1# type: ignore
2import os
3import shutil
4import sys
5import tempfile
6import traceback
7from datetime import datetime, timedelta
8from typing import Any
9
10import httpx
11import numpy as np
12import pytest
13
14import chromadb
15import chromadb.server.fastapi
16from chromadb.api.fastapi import FastAPI
17from chromadb.api.types import (
18 Document,
19 EmbeddingFunction,
20 QueryResult,
21 TYPE_KEY,
22 SPARSE_VECTOR_TYPE_VALUE,
23)
24from chromadb.config import Settings
25from chromadb.errors import (
26 ChromaError,
27 NotFoundError,
28 InvalidArgumentError,
29)
30from chromadb.utils.embedding_functions import DefaultEmbeddingFunction
31
32
33@pytest.fixture
34def persist_dir():
35 return tempfile.mkdtemp()
36
37
38@pytest.fixture
39def local_persist_api(persist_dir):
40 client = chromadb.Client(
41 Settings(
42 chroma_api_impl="chromadb.api.segment.SegmentAPI",
43 chroma_sysdb_impl="chromadb.db.impl.sqlite.SqliteDB",
44 chroma_producer_impl="chromadb.db.impl.sqlite.SqliteDB",
45 chroma_consumer_impl="chromadb.db.impl.sqlite.SqliteDB",
46 chroma_segment_manager_impl="chromadb.segment.impl.manager.local.LocalSegmentManager",
47 allow_reset=True,
48 is_persistent=True,
49 persist_directory=persist_dir,
50 ),
51 )
52 yield client
53 client.clear_system_cache()
54 if os.path.exists(persist_dir):
55 shutil.rmtree(persist_dir, ignore_errors=True)
56
57
58# https://docs.pytest.org/en/6.2.x/fixture.html#fixtures-can-be-requested-more-than-once-per-test-return-values-are-cached
59@pytest.fixture
60def local_persist_api_cache_bust(persist_dir):
61 client = chromadb.Client(
62 Settings(
63 chroma_api_impl="chromadb.api.segment.SegmentAPI",
64 chroma_sysdb_impl="chromadb.db.impl.sqlite.SqliteDB",
65 chroma_producer_impl="chromadb.db.impl.sqlite.SqliteDB",
66 chroma_consumer_impl="chromadb.db.impl.sqlite.SqliteDB",
67 chroma_segment_manager_impl="chromadb.segment.impl.manager.local.LocalSegmentManager",
68 allow_reset=True,
69 is_persistent=True,
70 persist_directory=persist_dir,
71 ),
72 )
73 yield client
74 client.clear_system_cache()
75 if os.path.exists(persist_dir):
76 shutil.rmtree(persist_dir, ignore_errors=True)
77
78
79def approx_equal(a, b, tolerance=1e-6) -> bool:
80 return abs(a - b) < tolerance
81
82
83def vector_approx_equal(a, b, tolerance: float = 1e-6) -> bool:
84 if len(a) != len(b):
85 return False
86 return all([approx_equal(a, b, tolerance) for a, b in zip(a, b)])
87
88
89@pytest.mark.parametrize("api_fixture", [local_persist_api])
90def test_persist_index_loading(api_fixture, request):
91 client = request.getfixturevalue("local_persist_api")
92 client.reset()
93 collection = client.create_collection("test")
94 collection.add(ids="id1", documents="hello")
95
96 api2 = request.getfixturevalue("local_persist_api_cache_bust")
97 collection = api2.get_collection("test")
98
99 includes = ["embeddings", "documents", "metadatas", "distances"]
100 nn = collection.query(
101 query_texts="hello",
102 n_results=1,
103 include=["embeddings", "documents", "metadatas", "distances"],
104 )
105 for key in nn.keys():
106 if (key in includes) or (key == "ids"):
107 assert len(nn[key]) == 1
108 elif key == "included":
109 assert set(nn[key]) == set(includes)
110 else:
111 assert nn[key] is None
112
113
114@pytest.mark.parametrize("api_fixture", [local_persist_api])
115def test_persist_index_loading_embedding_function(api_fixture, request):
116 class TestEF(EmbeddingFunction[Document]):
117 def __call__(self, input):
118 return [np.array([1, 2, 3]) for _ in range(len(input))]
119
120 def __init__(self, *args: Any, **kwargs: Any) -> None:
121 super().__init__(*args, **kwargs)
122
123 def name(self) -> str:
124 return "test"
125
126 def build_from_config(self, config: dict[str, Any]) -> None:
127 pass
128
129 def get_config(self) -> dict[str, Any]:
130 return {}
131
132 client = request.getfixturevalue("local_persist_api")
133 client.reset()
134 collection = client.create_collection("test", embedding_function=TestEF())
135 collection.add(ids="id1", documents="hello")
136
137 client2 = request.getfixturevalue("local_persist_api_cache_bust")
138 collection = client2.get_collection("test", embedding_function=TestEF())
139
140 includes = ["embeddings", "documents", "metadatas", "distances"]
141 nn = collection.query(
142 query_texts="hello",
143 n_results=1,
144 include=includes,
145 )
146 for key in nn.keys():
147 if (key in includes) or (key == "ids"):
148 assert len(nn[key]) == 1
149 elif key == "included":
150 assert set(nn[key]) == set(includes)
151 else:
152 assert nn[key] is None
153
154
155@pytest.mark.parametrize("api_fixture", [local_persist_api])
156def test_persist_index_get_or_create_embedding_function(api_fixture, request):
157 class TestEF(EmbeddingFunction[Document]):
158 def __call__(self, input):
159 return [np.array([1, 2, 3]) for _ in range(len(input))]
160
161 def __init__(self, *args: Any, **kwargs: Any) -> None:
162 super().__init__(*args, **kwargs)
163
164 def name(self) -> str:
165 return "test"
166
167 def build_from_config(self, config: dict[str, Any]) -> None:
168 pass
169
170 def get_config(self) -> dict[str, Any]:
171 return {}
172
173 api = request.getfixturevalue("local_persist_api")
174 api.reset()
175 collection = api.get_or_create_collection("test", embedding_function=TestEF())
176 collection.add(ids="id1", documents="hello")
177
178 api2 = request.getfixturevalue("local_persist_api_cache_bust")
179 collection = api2.get_or_create_collection("test", embedding_function=TestEF())
180
181 includes = ["embeddings", "documents", "metadatas", "distances"]
182 nn = collection.query(
183 query_texts="hello",
184 n_results=1,
185 include=includes,
186 )
187
188 for key in nn.keys():
189 if (key in includes) or (key == "ids"):
190 assert len(nn[key]) == 1
191 elif key == "included":
192 assert set(nn[key]) == set(includes)
193 else:
194 assert nn[key] is None
195
196 assert nn["ids"] == [["id1"]]
197 assert nn["embeddings"][0][0].tolist() == [1, 2, 3]
198 assert nn["documents"] == [["hello"]]
199 assert nn["distances"] == [[0]]
200
201
202@pytest.mark.parametrize("api_fixture", [local_persist_api])
203def test_persist(api_fixture, request):
204 client = request.getfixturevalue(api_fixture.__name__)
205
206 client.reset()
207
208 collection = client.create_collection("testspace")
209
210 collection.add(**batch_records)
211
212 assert collection.count() == 2
213
214 client = request.getfixturevalue(api_fixture.__name__)
215 collection = client.get_collection("testspace")
216 assert collection.count() == 2
217
218 client.delete_collection("testspace")
219
220 client = request.getfixturevalue(api_fixture.__name__)
221 assert client.list_collections() == []
222
223
224def test_heartbeat(client):
225 heartbeat_ns = client.heartbeat()
226 assert isinstance(heartbeat_ns, int)
227
228 heartbeat_s = heartbeat_ns // 10**9
229 heartbeat = datetime.fromtimestamp(heartbeat_s)
230 assert heartbeat > datetime.now() - timedelta(seconds=10)
231
232
233def test_max_batch_size(client):
234 batch_size = client.get_max_batch_size()
235 assert batch_size > 0
236
237
238def test_supports_base64_encoding(client):
239 if not isinstance(client, FastAPI):
240 pytest.skip("Not a FastAPI instance")
241
242 client.reset()
243
244 supports_base64_encoding = client.supports_base64_encoding()
245 assert supports_base64_encoding is True
246
247
248def test_supports_base64_encoding_legacy(client):
249 if not isinstance(client, FastAPI):
250 pytest.skip("Not a FastAPI instance")
251
252 client.reset()
253
254 # legacy server does not give back supports_base64_encoding
255 client.pre_flight_checks = {
256 "max_batch_size": 100,
257 }
258
259 assert client.supports_base64_encoding() is False
260 assert client.get_max_batch_size() == 100
261
262
263def test_pre_flight_checks(client):
264 if not isinstance(client, FastAPI):
265 pytest.skip("Not a FastAPI instance")
266
267 resp = httpx.get(f"{client._api_url}/pre-flight-checks")
268 assert resp.status_code == 200
269 assert resp.json() is not None
270 assert "max_batch_size" in resp.json().keys()
271 assert "supports_base64_encoding" in resp.json().keys()
272
273
274batch_records = {
275 "embeddings": [[1.1, 2.3, 3.2], [1.2, 2.24, 3.2]],
276 "ids": ["https://example.com/1", "https://example.com/2"],
277}
278
279
280def test_add(client):
281 client.reset()
282
283 collection = client.create_collection("testspace")
284
285 collection.add(**batch_records)
286
287 assert collection.count() == 2
288
289
290def test_collection_add_with_invalid_collection_throws(client):
291 client.reset()
292 collection = client.create_collection("test")
293 client.delete_collection("test")
294
295 with pytest.raises(NotFoundError, match=r"Collection .* does not exist"):
296 collection.add(**batch_records)
297
298
299def test_get_or_create(client):
300 client.reset()
301
302 collection = client.create_collection("testspace")
303
304 collection.add(**batch_records)
305
306 assert collection.count() == 2
307
308 with pytest.raises(Exception):
309 collection = client.create_collection("testspace")
310
311 collection = client.get_or_create_collection("testspace")
312
313 assert collection.count() == 2
314
315
316minimal_records = {
317 "embeddings": [[1.1, 2.3, 3.2], [1.2, 2.24, 3.2]],
318 "ids": ["https://example.com/1", "https://example.com/2"],
319}
320
321
322def test_add_minimal(client):
323 client.reset()
324
325 collection = client.create_collection("testspace")
326
327 collection.add(**minimal_records)
328
329 assert collection.count() == 2
330
331
332def test_get_from_db(client):
333 client.reset()
334 collection = client.create_collection("testspace")
335 collection.add(**batch_records)
336 includes = ["embeddings", "documents", "metadatas"]
337 records = collection.get(include=includes)
338 for key in records.keys():
339 if (key in includes) or (key == "ids"):
340 assert len(records[key]) == 2
341 elif key == "included":
342 assert set(records[key]) == set(includes)
343 else:
344 assert records[key] is None
345
346
347def test_collection_get_with_invalid_collection_throws(client):
348 client.reset()
349 collection = client.create_collection("test")
350 client.delete_collection("test")
351
352 with pytest.raises(NotFoundError, match=r"Collection .* does not exist"):
353 collection.get()
354
355
356def test_reset_db(client):
357 client.reset()
358
359 collection = client.create_collection("testspace")
360 collection.add(**batch_records)
361 assert collection.count() == 2
362
363 client.reset()
364 assert len(client.list_collections()) == 0
365
366
367def test_get_nearest_neighbors(client):
368 client.reset()
369 collection = client.create_collection("testspace")
370 collection.add(**batch_records)
371
372 includes = ["embeddings", "documents", "metadatas", "distances"]
373 nn = collection.query(
374 query_embeddings=[1.1, 2.3, 3.2],
375 n_results=1,
376 include=includes,
377 )
378 for key in nn.keys():
379 if (key in includes) or (key == "ids"):
380 assert len(nn[key]) == 1
381 elif key == "included":
382 assert set(nn[key]) == set(includes)
383 else:
384 assert nn[key] is None
385
386 nn = collection.query(
387 query_embeddings=[[1.1, 2.3, 3.2]],
388 n_results=1,
389 include=includes,
390 )
391 for key in nn.keys():
392 if (key in includes) or (key == "ids"):
393 assert len(nn[key]) == 1
394 elif key == "included":
395 assert set(nn[key]) == set(includes)
396 else:
397 assert nn[key] is None
398
399 nn = collection.query(
400 query_embeddings=[[1.1, 2.3, 3.2], [0.1, 2.3, 4.5]],
401 n_results=1,
402 include=includes,
403 )
404 for key in nn.keys():
405 if (key in includes) or (key == "ids"):
406 assert len(nn[key]) == 2
407 elif key == "included":
408 assert set(nn[key]) == set(includes)
409 else:
410 assert nn[key] is None
411
412
413def test_delete(client):
414 client.reset()
415 collection = client.create_collection("testspace")
416 collection.add(**batch_records)
417 assert collection.count() == 2
418
419 with pytest.raises(Exception):
420 collection.delete()
421
422
423def test_delete_returns_delete_result(client):
424 client.reset()
425 collection = client.create_collection("testspace")
426 collection.add(**batch_records)
427 assert collection.count() == 2
428 result = collection.delete(ids=batch_records["ids"])
429 assert isinstance(result, dict)
430 assert "deleted" in result
431 assert result["deleted"] >= 0
432
433
434def test_delete_with_limit(client):
435 client.reset()
436 collection = client.create_collection(
437 "testspace",
438 metadata={"hnsw:space": "l2"},
439 )
440 collection.add(
441 ids=["id1", "id2", "id3", "id4", "id5"],
442 embeddings=[[1, 0, 0], [0, 1, 0], [0, 0, 1], [1, 1, 0], [0, 1, 1]],
443 metadatas=[
444 {"category": "a"},
445 {"category": "a"},
446 {"category": "a"},
447 {"category": "b"},
448 {"category": "b"},
449 ],
450 )
451 assert collection.count() == 5
452
453 # Delete at most 2 records matching category=a (there are 3 total)
454 result = collection.delete(where={"category": "a"}, limit=2)
455 assert result["deleted"] == 2
456 assert collection.count() == 3
457
458
459def test_delete_with_limit_zero_is_noop(client):
460 client.reset()
461 collection = client.create_collection(
462 "testspace",
463 metadata={"hnsw:space": "l2"},
464 )
465 collection.add(
466 ids=["id1", "id2"],
467 embeddings=[[1, 0, 0], [0, 1, 0]],
468 metadatas=[{"category": "a"}, {"category": "a"}],
469 )
470 assert collection.count() == 2
471 result = collection.delete(where={"category": "a"}, limit=0)
472 assert result["deleted"] == 0
473 assert collection.count() == 2
474
475
476def test_delete_with_limit_requires_where(client):
477 client.reset()
478 collection = client.create_collection("testspace")
479 collection.add(**batch_records)
480 with pytest.raises(ValueError, match="limit can only be specified"):
481 collection.delete(ids=["id1"], limit=5)
482
483
484def test_delete_with_index(client):
485 client.reset()
486 collection = client.create_collection("testspace")
487 collection.add(**batch_records)
488 assert collection.count() == 2
489 collection.query(query_embeddings=[[1.1, 2.3, 3.2]], n_results=1)
490
491
492def test_collection_delete_with_invalid_collection_throws(client):
493 client.reset()
494 collection = client.create_collection("test")
495 client.delete_collection("test")
496
497 with pytest.raises(NotFoundError, match=r"Collection .* does not exist"):
498 collection.delete(ids=["id1"])
499
500
501def test_count(client):
502 client.reset()
503 collection = client.create_collection("testspace")
504 assert collection.count() == 0
505 collection.add(**batch_records)
506 assert collection.count() == 2
507
508
509def test_collection_count_with_invalid_collection_throws(client):
510 client.reset()
511 collection = client.create_collection("test")
512 client.delete_collection("test")
513
514 with pytest.raises(NotFoundError, match=r"Collection .* does not exist"):
515 collection.count()
516
517
518def test_modify(client):
519 client.reset()
520 collection = client.create_collection("testspace")
521 collection.modify(name="testspace2")
522
523 # collection name is modify
524 assert collection.name == "testspace2"
525
526
527def test_collection_modify_with_invalid_collection_throws(client):
528 client.reset()
529 collection = client.create_collection("test")
530 client.delete_collection("test")
531
532 with pytest.raises(NotFoundError, match=r"Collection .* does not exist"):
533 collection.modify(name="test2")
534
535
536def test_modify_error_on_existing_name(client):
537 client.reset()
538
539 client.create_collection("testspace")
540 c2 = client.create_collection("testspace2")
541
542 with pytest.raises(Exception):
543 c2.modify(name="testspace")
544
545
546def test_modify_warn_on_DF_change(client, caplog):
547 client.reset()
548
549 collection = client.create_collection("testspace")
550
551 with pytest.raises(Exception, match="not supported"):
552 collection.modify(metadata={"hnsw:space": "cosine"})
553
554
555def test_metadata_cru(client):
556 client.reset()
557 metadata_a = {"a": 1, "b": 2}
558 # Test create metadata
559 collection = client.create_collection("testspace", metadata=metadata_a)
560 assert collection.metadata is not None
561 assert collection.metadata["a"] == 1
562 assert collection.metadata["b"] == 2
563
564 # Test get metadata
565 collection = client.get_collection("testspace")
566 assert collection.metadata is not None
567 assert collection.metadata["a"] == 1
568 assert collection.metadata["b"] == 2
569
570 # Test modify metadata
571 collection.modify(metadata={"a": 2, "c": 3})
572 assert collection.metadata["a"] == 2
573 assert collection.metadata["c"] == 3
574 assert "b" not in collection.metadata
575
576 # Test get after modify metadata
577 collection = client.get_collection("testspace")
578 assert collection.metadata is not None
579 assert collection.metadata["a"] == 2
580 assert collection.metadata["c"] == 3
581 assert "b" not in collection.metadata
582
583 # Test name exists get_or_create_metadata
584 collection = client.get_or_create_collection("testspace")
585 assert collection.metadata is not None
586 assert collection.metadata["a"] == 2
587 assert collection.metadata["c"] == 3
588
589 # Test name exists create metadata
590 collection = client.get_or_create_collection("testspace2")
591 assert collection.metadata is None
592
593 # Test list collections
594 collections = client.list_collections()
595 for collection in collections:
596 if collection.name == "testspace":
597 assert collection.metadata is not None
598 assert collection.metadata["a"] == 2
599 assert collection.metadata["c"] == 3
600 elif collection.name == "testspace2":
601 assert collection.metadata is None
602
603
604def test_increment_index_on(client):
605 client.reset()
606 collection = client.create_collection("testspace")
607 collection.add(**batch_records)
608 assert collection.count() == 2
609
610 includes = ["embeddings", "documents", "metadatas", "distances"]
611 # increment index
612 nn = collection.query(
613 query_embeddings=[[1.1, 2.3, 3.2]],
614 n_results=1,
615 include=includes,
616 )
617 for key in nn.keys():
618 if (key in includes) or (key == "ids"):
619 assert len(nn[key]) == 1
620 elif key == "included":
621 assert set(nn[key]) == set(includes)
622 else:
623 assert nn[key] is None
624
625
626def test_add_a_collection(client):
627 client.reset()
628 client.create_collection("testspace")
629
630 # get collection does not throw an error
631 collection = client.get_collection("testspace")
632 assert collection.name == "testspace"
633
634 # get collection should throw an error if collection does not exist
635 with pytest.raises(Exception):
636 collection = client.get_collection("testspace2")
637
638
639def test_get_collection_by_id(client):
640 import uuid
641
642 client.reset()
643 collection = client.create_collection("testspace", metadata={"key": "value"})
644 collection_id = collection.id
645
646 retrieved = client.get_collection_by_id(collection_id)
647 assert retrieved.name == "testspace"
648 assert retrieved.id == collection_id
649 assert retrieved.metadata == {"key": "value"}
650
651 with pytest.raises(NotFoundError):
652 client.get_collection_by_id(uuid.uuid4())
653
654
655def test_error_includes_trace_id(http_client):
656 http_client.reset()
657
658 with pytest.raises(ChromaError) as error:
659 http_client.get_collection("testspace2")
660
661 assert error.value.trace_id is not None
662
663
664def test_list_collections(client):
665 client.reset()
666 client.create_collection("testspace")
667 client.create_collection("testspace2")
668
669 # get collection does not throw an error
670 collections = client.list_collections()
671 assert len(collections) == 2
672
673
674def test_reset(client):
675 client.reset()
676 client.create_collection("testspace")
677 client.create_collection("testspace2")
678
679 # get collection does not throw an error
680 collections = client.list_collections()
681 assert len(collections) == 2
682
683 client.reset()
684 collections = client.list_collections()
685 assert len(collections) == 0
686
687
688def test_peek(client):
689 client.reset()
690 collection = client.create_collection("testspace")
691 collection.add(**batch_records)
692 assert collection.count() == 2
693
694 # peek
695 peek = collection.peek()
696 print(peek)
697 for key in peek.keys():
698 if key in ["embeddings", "documents", "metadatas"] or key == "ids":
699 assert len(peek[key]) == 2
700 elif key == "included":
701 assert set(peek[key]) == set(["embeddings", "metadatas", "documents"])
702 else:
703 assert peek[key] is None
704
705
706def test_collection_peek_with_invalid_collection_throws(client):
707 client.reset()
708 collection = client.create_collection("test")
709 client.delete_collection("test")
710
711 with pytest.raises(NotFoundError, match=r"Collection .* does not exist"):
712 collection.peek()
713
714
715def test_collection_query_with_invalid_collection_throws(client):
716 client.reset()
717 collection = client.create_collection("test")
718 client.delete_collection("test")
719
720 with pytest.raises(NotFoundError, match=r"Collection .* does not exist"):
721 collection.query(query_texts=["test"])
722
723
724def test_collection_update_with_invalid_collection_throws(client):
725 client.reset()
726 collection = client.create_collection("test")
727 client.delete_collection("test")
728
729 with pytest.raises(NotFoundError, match=r"Collection .* does not exist"):
730 collection.update(ids=["id1"], documents=["test"])
731
732
733# TEST METADATA AND METADATA FILTERING
734# region
735
736metadata_records = {
737 "embeddings": [[1.1, 2.3, 3.2], [1.2, 2.24, 3.2]],
738 "ids": ["id1", "id2"],
739 "metadatas": [
740 {"int_value": 1, "string_value": "one", "float_value": 1.001},
741 {"int_value": 2},
742 ],
743}
744
745
746def test_metadata_add_get_int_float(client):
747 client.reset()
748 collection = client.create_collection("test_int")
749 collection.add(**metadata_records)
750
751 items = collection.get(ids=["id1", "id2"])
752 assert items["metadatas"][0]["int_value"] == 1
753 assert items["metadatas"][0]["float_value"] == 1.001
754 assert items["metadatas"][1]["int_value"] == 2
755 assert isinstance(items["metadatas"][0]["int_value"], int)
756 assert isinstance(items["metadatas"][0]["float_value"], float)
757
758
759def test_metadata_add_query_int_float(client):
760 client.reset()
761 collection = client.create_collection("test_int")
762 collection.add(**metadata_records)
763
764 items: QueryResult = collection.query(
765 query_embeddings=[[1.1, 2.3, 3.2]], n_results=1
766 )
767 assert items["metadatas"] is not None
768 assert items["metadatas"][0][0]["int_value"] == 1
769 assert items["metadatas"][0][0]["float_value"] == 1.001
770 assert isinstance(items["metadatas"][0][0]["int_value"], int)
771 assert isinstance(items["metadatas"][0][0]["float_value"], float)
772
773
774def test_metadata_get_where_string(client):
775 client.reset()
776 collection = client.create_collection("test_int")
777 collection.add(**metadata_records)
778
779 items = collection.get(where={"string_value": "one"})
780 assert items["metadatas"][0]["int_value"] == 1
781 assert items["metadatas"][0]["string_value"] == "one"
782
783
784def test_metadata_get_where_int(client):
785 client.reset()
786 collection = client.create_collection("test_int")
787 collection.add(**metadata_records)
788
789 items = collection.get(where={"int_value": 1})
790 assert items["metadatas"][0]["int_value"] == 1
791 assert items["metadatas"][0]["string_value"] == "one"
792
793
794def test_metadata_get_where_float(client):
795 client.reset()
796 collection = client.create_collection("test_int")
797 collection.add(**metadata_records)
798
799 items = collection.get(where={"float_value": 1.001})
800 assert items["metadatas"][0]["int_value"] == 1
801 assert items["metadatas"][0]["string_value"] == "one"
802 assert items["metadatas"][0]["float_value"] == 1.001
803
804
805def test_metadata_update_get_int_float(client):
806 client.reset()
807 collection = client.create_collection("test_int")
808 collection.add(**metadata_records)
809
810 collection.update(
811 ids=["id1"],
812 metadatas=[{"int_value": 2, "string_value": "two", "float_value": 2.002}],
813 )
814 items = collection.get(ids=["id1"])
815 assert items["metadatas"][0]["int_value"] == 2
816 assert items["metadatas"][0]["string_value"] == "two"
817 assert items["metadatas"][0]["float_value"] == 2.002
818
819
820bad_metadata_records = {
821 "embeddings": [[1.1, 2.3, 3.2], [1.2, 2.24, 3.2]],
822 "ids": ["id1", "id2"],
823 "metadatas": [{"value": {"nested": "5"}}, {"value": [1, 2, 3]}],
824}
825
826
827def test_metadata_validation_add(client):
828 client.reset()
829 collection = client.create_collection("test_metadata_validation")
830 with pytest.raises(ValueError, match="metadata"):
831 collection.add(**bad_metadata_records)
832
833
834def test_metadata_validation_update(client):
835 client.reset()
836 collection = client.create_collection("test_metadata_validation")
837 collection.add(**metadata_records)
838 with pytest.raises(ValueError, match="metadata"):
839 collection.update(ids=["id1"], metadatas={"value": {"nested": "5"}})
840
841
842def test_where_validation_get(client):
843 client.reset()
844 collection = client.create_collection("test_where_validation")
845 with pytest.raises(ValueError, match="where"):
846 collection.get(where={"value": {"nested": "5"}})
847
848
849def test_where_validation_query(client):
850 client.reset()
851 collection = client.create_collection("test_where_validation")
852 with pytest.raises(ValueError, match="where"):
853 collection.query(query_embeddings=[0, 0, 0], where={"value": {"nested": "5"}})
854
855
856operator_records = {
857 "embeddings": [[1.1, 2.3, 3.2], [1.2, 2.24, 3.2]],
858 "ids": ["id1", "id2"],
859 "metadatas": [
860 {"int_value": 1, "string_value": "one", "float_value": 1.001},
861 {"int_value": 2, "float_value": 2.002, "string_value": "two"},
862 ],
863}
864
865
866def test_where_lt(client):
867 client.reset()
868 collection = client.create_collection("test_where_lt")
869 collection.add(**operator_records)
870 items = collection.get(where={"int_value": {"$lt": 2}})
871 assert len(items["metadatas"]) == 1
872
873
874def test_where_lte(client):
875 client.reset()
876 collection = client.create_collection("test_where_lte")
877 collection.add(**operator_records)
878 items = collection.get(where={"int_value": {"$lte": 2.0}})
879 assert len(items["metadatas"]) == 2
880
881
882def test_where_gt(client):
883 client.reset()
884 collection = client.create_collection("test_where_lte")
885 collection.add(**operator_records)
886 items = collection.get(where={"float_value": {"$gt": -1.4}})
887 assert len(items["metadatas"]) == 2
888
889
890def test_where_gte(client):
891 client.reset()
892 collection = client.create_collection("test_where_lte")
893 collection.add(**operator_records)
894 items = collection.get(where={"float_value": {"$gte": 2.002}})
895 assert len(items["metadatas"]) == 1
896
897
898def test_where_ne_string(client):
899 client.reset()
900 collection = client.create_collection("test_where_lte")
901 collection.add(**operator_records)
902 items = collection.get(where={"string_value": {"$ne": "two"}})
903 assert len(items["metadatas"]) == 1
904
905
906def test_where_ne_eq_number(client):
907 client.reset()
908 collection = client.create_collection("test_where_lte")
909 collection.add(**operator_records)
910 items = collection.get(where={"int_value": {"$ne": 1}})
911 assert len(items["metadatas"]) == 1
912 items = collection.get(where={"float_value": {"$eq": 2.002}})
913 assert len(items["metadatas"]) == 1
914
915
916def test_where_valid_operators(client):
917 client.reset()
918 collection = client.create_collection("test_where_valid_operators")
919 collection.add(**operator_records)
920 with pytest.raises(ValueError):
921 collection.get(where={"int_value": {"$invalid": 2}})
922
923 with pytest.raises(ValueError):
924 collection.get(where={"int_value": {"$lt": "2"}})
925
926 with pytest.raises(ValueError):
927 collection.get(where={"int_value": {"$lt": 2, "$gt": 1}})
928
929 # Test invalid $and, $or
930 with pytest.raises(ValueError):
931 collection.get(where={"$and": {"int_value": {"$lt": 2}}})
932
933 with pytest.raises(ValueError):
934 collection.get(
935 where={"int_value": {"$lt": 2}, "$or": {"int_value": {"$gt": 1}}}
936 )
937
938 with pytest.raises(ValueError):
939 collection.get(
940 where={"$gt": [{"int_value": {"$lt": 2}}, {"int_value": {"$gt": 1}}]}
941 )
942
943 with pytest.raises(ValueError):
944 collection.get(where={"$or": [{"int_value": {"$lt": 2}}]})
945
946 with pytest.raises(ValueError):
947 collection.get(where={"$or": []})
948
949 # $contains on a metadata key is now valid (array contains).
950 # But a bare $contains at the top level (no field key) is still invalid.
951 with pytest.raises(ValueError):
952 collection.get(where={"$contains": "test"})
953
954
955# TODO: Define the dimensionality of these embeddingds in terms of the default record
956bad_dimensionality_records = {
957 "embeddings": [[1.1, 2.3, 3.2, 4.5], [1.2, 2.24, 3.2, 4.5]],
958 "ids": ["id1", "id2"],
959}
960
961bad_dimensionality_query = {
962 "query_embeddings": [[1.1, 2.3, 3.2, 4.5], [1.2, 2.24, 3.2, 4.5]],
963}
964
965bad_number_of_results_query = {
966 "query_embeddings": [[1.1, 2.3, 3.2], [1.2, 2.24, 3.2]],
967 "n_results": 100,
968}
969
970
971def test_dimensionality_validation_add(client):
972 client.reset()
973 collection = client.create_collection("test_dimensionality_validation")
974 collection.add(**minimal_records)
975
976 with pytest.raises(Exception) as e:
977 collection.add(**bad_dimensionality_records)
978 assert "dimension" in str(e.value)
979
980
981def test_dimensionality_validation_query(client):
982 client.reset()
983 collection = client.create_collection("test_dimensionality_validation_query")
984 collection.add(**minimal_records)
985
986 with pytest.raises(Exception) as e:
987 collection.query(**bad_dimensionality_query)
988 assert "dimension" in str(e.value)
989
990
991def test_query_document_valid_operators(client):
992 client.reset()
993 collection = client.create_collection("test_where_valid_operators")
994 collection.add(**operator_records)
995 with pytest.raises(ValueError, match="where document"):
996 collection.get(where_document={"$lt": {"$nested": 2}})
997
998 with pytest.raises(ValueError, match="where document"):
999 collection.query(query_embeddings=[0, 0, 0], where_document={"$contains": 2})
1000
1001 with pytest.raises(ValueError, match="where document"):
1002 collection.get(where_document={"$contains": []})
1003
1004 # Test invalid $contains
1005 with pytest.raises(ValueError, match="where document"):
1006 collection.get(where_document={"$contains": {"text": "hello"}})
1007
1008 # Test invalid $not_contains
1009 with pytest.raises(ValueError, match="where document"):
1010 collection.get(where_document={"$not_contains": {"text": "hello"}})
1011
1012 # Test invalid $and, $or
1013 with pytest.raises(ValueError):
1014 collection.get(where_document={"$and": {"$unsupported": "doc"}})
1015
1016 with pytest.raises(ValueError):
1017 collection.get(
1018 where_document={"$or": [{"$unsupported": "doc"}, {"$unsupported": "doc"}]}
1019 )
1020
1021 with pytest.raises(ValueError):
1022 collection.get(where_document={"$or": [{"$contains": "doc"}]})
1023
1024 with pytest.raises(ValueError):
1025 collection.get(where_document={"$or": []})
1026
1027 with pytest.raises(ValueError):
1028 collection.get(
1029 where_document={
1030 "$or": [{"$and": [{"$contains": "doc"}]}, {"$contains": "doc"}]
1031 }
1032 )
1033
1034
1035contains_records = {
1036 "embeddings": [[1.1, 2.3, 3.2], [1.2, 2.24, 3.2]],
1037 "documents": ["this is doc1 and it's great!", "doc2 is also great!"],
1038 "ids": ["id1", "id2"],
1039 "metadatas": [
1040 {"int_value": 1, "string_value": "one", "float_value": 1.001},
1041 {"int_value": 2, "float_value": 2.002, "string_value": "two"},
1042 ],
1043}
1044
1045
1046def test_get_where_document(client):
1047 client.reset()
1048 collection = client.create_collection("test_get_where_document")
1049 collection.add(**contains_records)
1050
1051 items = collection.get(where_document={"$contains": "doc1"})
1052 assert len(items["metadatas"]) == 1
1053
1054 items = collection.get(where_document={"$contains": "great"})
1055 assert len(items["metadatas"]) == 2
1056
1057 items = collection.get(where_document={"$contains": "bad"})
1058 assert len(items["metadatas"]) == 0
1059
1060
1061def test_query_where_document(client):
1062 client.reset()
1063 collection = client.create_collection("test_query_where_document")
1064 collection.add(**contains_records)
1065
1066 items = collection.query(
1067 query_embeddings=[1, 0, 0], where_document={"$contains": "doc1"}, n_results=1
1068 )
1069 assert len(items["metadatas"][0]) == 1
1070
1071 items = collection.query(
1072 query_embeddings=[0, 0, 0], where_document={"$contains": "great"}, n_results=2
1073 )
1074 assert len(items["metadatas"][0]) == 2
1075
1076 with pytest.raises(Exception) as e:
1077 items = collection.query(
1078 query_embeddings=[0, 0, 0], where_document={"$contains": "bad"}, n_results=1
1079 )
1080 assert "datapoints" in str(e.value)
1081
1082
1083def test_delete_where_document(client):
1084 client.reset()
1085 collection = client.create_collection("test_delete_where_document")
1086 collection.add(**contains_records)
1087
1088 collection.delete(where_document={"$contains": "doc1"})
1089 assert collection.count() == 1
1090
1091 collection.delete(where_document={"$contains": "bad"})
1092 assert collection.count() == 1
1093
1094 collection.delete(where_document={"$contains": "great"})
1095 assert collection.count() == 0
1096
1097
1098logical_operator_records = {
1099 "embeddings": [
1100 [1.1, 2.3, 3.2],
1101 [1.2, 2.24, 3.2],
1102 [1.3, 2.25, 3.2],
1103 [1.4, 2.26, 3.2],
1104 ],
1105 "ids": ["id1", "id2", "id3", "id4"],
1106 "metadatas": [
1107 {"int_value": 1, "string_value": "one", "float_value": 1.001, "is": "doc"},
1108 {"int_value": 2, "float_value": 2.002, "string_value": "two", "is": "doc"},
1109 {"int_value": 3, "float_value": 3.003, "string_value": "three", "is": "doc"},
1110 {"int_value": 4, "float_value": 4.004, "string_value": "four", "is": "doc"},
1111 ],
1112 "documents": [
1113 "this document is first and great",
1114 "this document is second and great",
1115 "this document is third and great",
1116 "this document is fourth and great",
1117 ],
1118}
1119
1120
1121def test_where_logical_operators(client):
1122 client.reset()
1123 collection = client.create_collection("test_logical_operators")
1124 collection.add(**logical_operator_records)
1125
1126 items = collection.get(
1127 where={
1128 "$and": [
1129 {"$or": [{"int_value": {"$gte": 3}}, {"float_value": {"$lt": 1.9}}]},
1130 {"is": "doc"},
1131 ]
1132 }
1133 )
1134 assert len(items["metadatas"]) == 3
1135
1136 items = collection.get(
1137 where={
1138 "$or": [
1139 {
1140 "$and": [
1141 {"int_value": {"$eq": 3}},
1142 {"string_value": {"$eq": "three"}},
1143 ]
1144 },
1145 {
1146 "$and": [
1147 {"int_value": {"$eq": 4}},
1148 {"string_value": {"$eq": "four"}},
1149 ]
1150 },
1151 ]
1152 }
1153 )
1154 assert len(items["metadatas"]) == 2
1155
1156 items = collection.get(
1157 where={
1158 "$and": [
1159 {
1160 "$or": [
1161 {"int_value": {"$eq": 1}},
1162 {"string_value": {"$eq": "two"}},
1163 ]
1164 },
1165 {
1166 "$or": [
1167 {"int_value": {"$eq": 2}},
1168 {"string_value": {"$eq": "one"}},
1169 ]
1170 },
1171 ]
1172 }
1173 )
1174 assert len(items["metadatas"]) == 2
1175
1176
1177def test_where_document_logical_operators(client):
1178 client.reset()
1179 collection = client.create_collection("test_document_logical_operators")
1180 collection.add(**logical_operator_records)
1181
1182 items = collection.get(
1183 where_document={
1184 "$and": [
1185 {"$contains": "first"},
1186 {"$contains": "doc"},
1187 ]
1188 }
1189 )
1190 assert len(items["metadatas"]) == 1
1191
1192 items = collection.get(
1193 where_document={
1194 "$or": [
1195 {"$contains": "first"},
1196 {"$contains": "second"},
1197 ]
1198 }
1199 )
1200 assert len(items["metadatas"]) == 2
