Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_api.py4024 linesDownload Raw Back to test
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

Showing the first 1,200 of 4024 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai