Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_schema_e2e.py2844 linesDownload Raw Back to api
1from chromadb.api import ClientAPI, ServerAPI
2from chromadb.api.types import (
3    Schema,
4    FtsIndexConfig,
5    SparseVectorIndexConfig,
6    SparseEmbeddingFunction,
7    SparseVector,
8    StringInvertedIndexConfig,
9    IntInvertedIndexConfig,
10    FloatInvertedIndexConfig,
11    BoolInvertedIndexConfig,
12    VectorIndexConfig,
13    SpannIndexConfig,
14    EmbeddingFunction,
15    Embeddings,
16)
17from chromadb.execution.expression.operator import Key
18from chromadb.test.conftest import (
19    ClientFactories,
20    is_spann_disabled_mode,
21    skip_if_not_cluster,
22    skip_reason_spann_disabled,
23    skip_reason_spann_enabled,
24)
25from chromadb.test.utils.wait_for_version_increase import (
26    get_collection_version,
27    wait_for_version_increase,
28)
29from chromadb.utils.embedding_functions import (
30    register_embedding_function,
31    register_sparse_embedding_function,
32)
33from chromadb.api.models.Collection import Collection
34from chromadb.api.models.CollectionCommon import CollectionCommon
35from chromadb.errors import InvalidArgumentError
36from chromadb.execution.expression import Knn, Search
37from chromadb.types import Collection as CollectionModel
38from typing import Any, Callable, Dict, List, Mapping, Optional, Tuple, cast
39from uuid import uuid4
40import numpy as np
41import pytest
42
43
44@register_embedding_function
45class SimpleEmbeddingFunction(EmbeddingFunction[List[str]]):
46    """Simple embedding function with stable configuration for persistence tests."""
47
48    def __init__(self, dim: int = 4):
49        self._dim = dim
50
51    def __call__(self, input: List[str]) -> Embeddings:
52        vector = [float(i) for i in range(self._dim)]
53        return cast(Embeddings, [vector for _ in input])
54
55    @staticmethod
56    def name() -> str:
57        return "simple_ef"
58
59    def get_config(self) -> Dict[str, Any]:
60        return {"dim": self._dim}
61
62    @staticmethod
63    def build_from_config(config: Dict[str, Any]) -> "SimpleEmbeddingFunction":
64        return SimpleEmbeddingFunction(dim=config["dim"])
65
66    def default_space(self) -> str:  # type: ignore[override]
67        return "cosine"
68
69
70@register_embedding_function
71class RecordingSearchEmbeddingFunction(EmbeddingFunction[List[str]]):
72    """Embedding function that records inputs for search embedding tests."""
73
74    def __init__(self, label: str = "default") -> None:
75        self._label = label
76        self.call_inputs: List[List[str]] = []
77        self.query_inputs: List[List[str]] = []
78
79    def __call__(self, input: List[str]) -> Embeddings:
80        self.call_inputs.append(list(input))
81        vectors = [[float(len(text)), float(len(text)) + 0.5] for text in input]
82        return cast(Embeddings, vectors)
83
84    def embed_query(self, input: List[str]) -> Embeddings:
85        self.query_inputs.append(list(input))
86        vectors = [[float(len(text)), float(len(text)) + 1.5] for text in input]
87        return cast(Embeddings, vectors)
88
89    @staticmethod
90    def name() -> str:
91        return "recording_search_ef"
92
93    def get_config(self) -> Dict[str, Any]:
94        return {"label": self._label}
95
96    @staticmethod
97    def build_from_config(config: Dict[str, Any]) -> "RecordingSearchEmbeddingFunction":
98        return RecordingSearchEmbeddingFunction(config.get("label", "default"))
99
100
101def test_schema_vector_config_persistence(
102    client_factories: "ClientFactories",
103) -> None:
104    """Ensure schema-provided SPANN settings persist across client restarts."""
105
106    client = client_factories.create_client_from_system()
107    client.reset()
108
109    collection_name = f"schema_spann_{uuid4().hex}"
110
111    schema = Schema()
112    schema.create_index(
113        config=VectorIndexConfig(
114            space="cosine",
115            spann=SpannIndexConfig(
116                search_nprobe=16,
117                write_nprobe=32,
118                ef_construction=120,
119                max_neighbors=24,
120            ),
121        )
122    )
123
124    collection = client.get_or_create_collection(
125        name=collection_name,
126        schema=schema,
127    )
128
129    persisted_schema = collection.schema
130    assert persisted_schema is not None
131
132    print(persisted_schema.serialize_to_json())
133
134    embedding_override = persisted_schema.keys["#embedding"].float_list
135    assert embedding_override is not None
136    vector_index = embedding_override.vector_index
137    assert vector_index is not None
138    assert vector_index.enabled is True
139    assert vector_index.config is not None
140    assert vector_index.config.space is not None
141    assert vector_index.config.space == "cosine"
142
143    client_reloaded = client_factories.create_client_from_system()
144    reloaded_collection = client_reloaded.get_collection(
145        name=collection_name,
146    )
147
148    reloaded_schema = reloaded_collection.schema
149    assert reloaded_schema is not None
150    reloaded_embedding_override = reloaded_schema.keys["#embedding"].float_list
151    assert reloaded_embedding_override is not None
152    reloaded_vector_index = reloaded_embedding_override.vector_index
153    assert reloaded_vector_index is not None
154    assert reloaded_vector_index.config is not None
155    assert reloaded_vector_index.config.space is not None
156    assert reloaded_vector_index.config.space == "cosine"
157
158
159def test_schema_vector_config_persistence_with_ef(
160    client_factories: "ClientFactories",
161) -> None:
162    """Ensure schema-provided SPANN settings persist across client restarts."""
163
164    client = client_factories.create_client_from_system()
165    client.reset()
166
167    collection_name = f"schema_spann_{uuid4().hex}"
168
169    schema = Schema()
170    embedding_function = SimpleEmbeddingFunction(dim=6)
171    schema.create_index(
172        config=VectorIndexConfig(
173            space="cosine",
174            embedding_function=embedding_function,
175            spann=SpannIndexConfig(
176                search_nprobe=16,
177                write_nprobe=32,
178                ef_construction=120,
179                max_neighbors=24,
180            ),
181        )
182    )
183
184    collection = client.get_or_create_collection(
185        name=collection_name,
186        schema=schema,
187    )
188
189    persisted_schema = collection.schema
190    assert persisted_schema is not None
191
192    print(persisted_schema.serialize_to_json())
193
194    embedding_override = persisted_schema.keys["#embedding"].float_list
195    assert embedding_override is not None
196    vector_index = embedding_override.vector_index
197    assert vector_index is not None
198    assert vector_index.enabled is True
199    assert vector_index.config is not None
200    assert vector_index.config.space is not None
201    assert vector_index.config.space == "cosine"
202
203    if not is_spann_disabled_mode:
204        assert vector_index.config.spann is not None
205        spann_config = vector_index.config.spann
206        assert spann_config.search_nprobe == 16
207        assert spann_config.write_nprobe == 32
208        assert spann_config.ef_construction == 120
209        assert spann_config.max_neighbors == 24
210    else:
211        assert vector_index.config.spann is None
212        assert vector_index.config.hnsw is not None
213        hnsw_config = vector_index.config.hnsw
214        assert hnsw_config.ef_construction == 100
215        assert hnsw_config.ef_search == 100
216        assert hnsw_config.max_neighbors == 16
217        assert hnsw_config.resize_factor == 1.2
218
219    ef = vector_index.config.embedding_function
220    assert ef is not None
221    assert ef.name() == "simple_ef"
222    assert ef.get_config() == {"dim": 6}
223
224    persisted_json = persisted_schema.serialize_to_json()
225    if not is_spann_disabled_mode:
226        spann_json = persisted_json["keys"]["#embedding"]["float_list"]["vector_index"][
227            "config"
228        ]["spann"]
229        assert spann_json["search_nprobe"] == 16
230        assert spann_json["write_nprobe"] == 32
231    else:
232        hnsw_json = persisted_json["keys"]["#embedding"]["float_list"]["vector_index"][
233            "config"
234        ]["hnsw"]
235        assert hnsw_json["ef_construction"] == 100
236        assert hnsw_json["ef_search"] == 100
237        assert hnsw_json["max_neighbors"] == 16
238
239    client_reloaded = client_factories.create_client_from_system()
240    reloaded_collection = client_reloaded.get_collection(
241        name=collection_name,
242    )
243
244    reloaded_schema = reloaded_collection.schema
245    assert reloaded_schema is not None
246    reloaded_embedding_override = reloaded_schema.keys["#embedding"].float_list
247    assert reloaded_embedding_override is not None
248    reloaded_vector_index = reloaded_embedding_override.vector_index
249    assert reloaded_vector_index is not None
250    assert reloaded_vector_index.config is not None
251    assert reloaded_vector_index.config.space is not None
252    assert reloaded_vector_index.config.space == "cosine"
253    if not is_spann_disabled_mode:
254        assert reloaded_vector_index.config.spann is not None
255        assert reloaded_vector_index.config.spann.search_nprobe == 16
256        assert reloaded_vector_index.config.spann.write_nprobe == 32
257    else:
258        assert reloaded_vector_index.config.hnsw is not None
259        assert reloaded_vector_index.config.hnsw.ef_construction == 100
260        assert reloaded_vector_index.config.hnsw.ef_search == 100
261        assert reloaded_vector_index.config.hnsw.max_neighbors == 16
262        assert reloaded_vector_index.config.hnsw.resize_factor == 1.2
263
264    config = reloaded_collection.configuration
265    assert config is not None
266    config_ef = config.get("embedding_function")
267    assert config_ef is not None
268    assert config_ef.name() == "simple_ef"
269    assert config_ef.get_config() == {"dim": 6}
270
271
272@register_sparse_embedding_function
273class DeterministicSparseEmbeddingFunction(SparseEmbeddingFunction[List[str]]):
274    """Sparse embedding function that emits predictable token/value pairs."""
275
276    def __init__(self, label: str = "det_sparse"):
277        self._label = label
278
279    def __call__(self, input: List[str]) -> List[SparseVector]:
280        return [
281            SparseVector(indices=[idx], values=[float(len(text) + idx)])
282            for idx, text in enumerate(input)
283        ]
284
285    @staticmethod
286    def name() -> str:
287        return "det_sparse"
288
289    def get_config(self) -> Dict[str, Any]:
290        return {"label": self._label}
291
292    @staticmethod
293    def build_from_config(
294        config: Dict[str, Any]
295    ) -> "DeterministicSparseEmbeddingFunction":
296        return DeterministicSparseEmbeddingFunction(config.get("label", "det_sparse"))
297
298
299def _create_isolated_collection(
300    client_factories: "ClientFactories",
301    schema: Optional[Schema] = None,
302    embedding_function: Optional[EmbeddingFunction[Any]] = None,
303) -> Tuple[Collection, ClientAPI]:
304    """Provision a new temporary collection and return it with the backing client."""
305    client = client_factories.create_client_from_system()
306    client.reset()
307
308    collection_name = f"schema_e2e_{uuid4().hex}"
309    if schema is not None:
310        collection = client.get_or_create_collection(
311            name=collection_name,
312            schema=schema,
313        )
314    else:
315        if embedding_function is not None:
316            collection = client.get_or_create_collection(
317                name=collection_name,
318                embedding_function=embedding_function,
319            )
320        else:
321            collection = client.get_or_create_collection(
322                name=collection_name,
323            )
324
325    return collection, client
326
327
328def _collect_knn_queries(rank: Any) -> List[Any]:
329    if rank is None:
330        return []
331
332    if isinstance(rank, Knn):
333        return [rank.query]
334
335    queries: List[Any] = []
336
337    child_rank = getattr(rank, "rank", None)
338    if child_rank is not None:
339        queries.extend(_collect_knn_queries(child_rank))
340
341    left_rank = getattr(rank, "left", None)
342    if left_rank is not None:
343        queries.extend(_collect_knn_queries(left_rank))
344
345    right_rank = getattr(rank, "right", None)
346    if right_rank is not None:
347        queries.extend(_collect_knn_queries(right_rank))
348
349    ranks = getattr(rank, "ranks", None)
350    if ranks:
351        for child in ranks:
352            queries.extend(_collect_knn_queries(child))
353
354    return queries
355
356
357def test_schema_defaults_enable_indexed_operations(
358    client_factories: "ClientFactories",
359) -> None:
360    """Validate default schema indexes support filtering, updates, and embeddings."""
361    collection, client = _create_isolated_collection(client_factories)
362
363    schema = collection.schema
364    assert schema is not None
365    assert schema.defaults is not None
366    assert schema.defaults.string is not None
367    string_index = schema.defaults.string.string_inverted_index
368    assert string_index is not None
369    assert string_index.enabled is True
370    assert schema.defaults.int_value is not None
371    int_index = schema.defaults.int_value.int_inverted_index
372    assert int_index is not None
373    assert int_index.enabled is True
374    assert schema.defaults.float_value is not None
375    float_index = schema.defaults.float_value.float_inverted_index
376    assert float_index is not None
377    assert float_index.enabled is True
378    assert schema.defaults.boolean is not None
379    bool_index = schema.defaults.boolean.bool_inverted_index
380    assert bool_index is not None
381    assert bool_index.enabled is True
382
383    document_override = schema.keys["#document"].string
384    assert document_override is not None
385    fts_index = document_override.fts_index
386    assert fts_index is not None
387    assert fts_index.enabled is True
388
389    embedding_override = schema.keys["#embedding"].float_list
390    assert embedding_override is not None
391    vector_index = embedding_override.vector_index
392    assert vector_index is not None
393    assert vector_index.enabled is True
394
395    ids = ["doc-1", "doc-2", "doc-3"]
396    documents = ["alpha", "beta", "gamma"]
397    metadatas: List[Mapping[str, Any]] = [
398        {"category": "news", "rating": 5, "price": 9.5, "is_active": True},
399        {"category": "science", "rating": 7, "price": 2.5, "is_active": False},
400        {"category": "news", "rating": 3, "price": 5.0, "is_active": True},
401    ]
402
403    collection.add(ids=ids, documents=documents, metadatas=metadatas)
404
405    filtered = collection.get(where={"category": "science"})
406    assert set(filtered["ids"]) == {"doc-2"}
407
408    numeric_filter = collection.get(where={"rating": 3})
409    assert set(numeric_filter["ids"]) == {"doc-3"}
410
411    bool_filter = collection.get(where={"is_active": False})
412    assert set(bool_filter["ids"]) == {"doc-2"}
413
414    collection.update(ids=["doc-1"], metadatas=[{"rating": 6, "category": "updates"}])
415    rating_after_update = collection.get(where={"rating": 6})
416    assert set(rating_after_update["ids"]) == {"doc-1"}
417
418    collection.upsert(
419        ids=["doc-2"],
420        documents=["beta-updated"],
421        metadatas=[{"price": 2.5, "category": "science"}],
422    )
423
424    embeddings_payload = collection.get(ids=["doc-1"], include=["embeddings"])
425    assert embeddings_payload["embeddings"] is not None
426    assert len(embeddings_payload["embeddings"]) == 1
427
428    # Ensure underlying schema persisted across fetches
429    reloaded = client.get_collection(collection.name)
430    assert reloaded.schema is not None
431    if not is_spann_disabled_mode:
432        assert reloaded.schema.serialize_to_json() == schema.serialize_to_json()
433
434
435def test_get_or_create_and_get_collection_preserve_schema(
436    client_factories: "ClientFactories",
437) -> None:
438    """Ensure repeated collection lookups reuse the persisted schema definition."""
439    base_schema = Schema()
440    base_schema.create_index(
441        key="custom_tag",
442        config=StringInvertedIndexConfig(),
443    )
444    base_schema.create_index(
445        key="importance",
446        config=IntInvertedIndexConfig(),
447    )
448
449    collection, client = _create_isolated_collection(
450        client_factories,
451        schema=base_schema,
452    )
453
454    assert collection.schema is not None
455    initial_schema_json = collection.schema.serialize_to_json()
456    assert "custom_tag" in initial_schema_json["keys"]
457    assert "importance" in initial_schema_json["keys"]
458
459    second_reference = client.get_or_create_collection(name=collection.name)
460    assert second_reference.schema is not None
461    assert second_reference.schema.serialize_to_json() == initial_schema_json
462
463    fetched = client.get_collection(name=collection.name)
464    assert fetched.schema is not None
465    assert fetched.schema.serialize_to_json() == initial_schema_json
466
467    second_reference.add(
468        ids=["schema-preserve"],
469        documents=["doc"],
470        metadatas=[{"custom_tag": "alpha", "importance": 10}],
471    )
472
473    stored = fetched.get(where={"custom_tag": "alpha"})
474    assert set(stored["ids"]) == {"schema-preserve"}
475
476
477def test_delete_collection_resets_schema_configuration(
478    client_factories: "ClientFactories",
479) -> None:
480    """Deleting and recreating a collection should drop prior schema overrides."""
481    schema = Schema()
482    schema.create_index(
483        key="transient_key",
484        config=StringInvertedIndexConfig(),
485    )
486
487    collection, client = _create_isolated_collection(
488        client_factories,
489        schema=schema,
490    )
491
492    assert collection.schema is not None
493    assert "transient_key" in collection.schema.keys
494
495    client.delete_collection(name=collection.name)
496
497    recreated = client.create_collection(name=collection.name)
498    assert recreated.schema is not None
499    recreated_json = recreated.schema.serialize_to_json()
500    baseline_json = Schema().serialize_to_json()
501    assert "transient_key" not in recreated_json["keys"]
502    assert set(recreated_json["keys"].keys()) == set(baseline_json["keys"].keys())
503
504
505@pytest.mark.skipif(not is_spann_disabled_mode, reason=skip_reason_spann_enabled)
506def test_sparse_vector_not_allowed_locally(
507    client_factories: "ClientFactories",
508) -> None:
509    """Sparse vector configs are not allowed to be created locally."""
510    schema = Schema()
511    schema.create_index(key="sparse_metadata", config=SparseVectorIndexConfig())
512    with pytest.raises(
513        InvalidArgumentError, match="Sparse vector indexing is not enabled in local"
514    ):
515        _create_isolated_collection(client_factories, schema=schema)
516
517
518@pytest.mark.skipif(is_spann_disabled_mode, reason=skip_reason_spann_disabled)
519def test_sparse_vector_source_key_and_index_constraints(
520    client_factories: "ClientFactories",
521) -> None:
522    """Sparse vector configs honor source key embedding and single-index enforcement."""
523    sparse_ef = DeterministicSparseEmbeddingFunction(label="source-test")
524
525    schema = Schema()
526    schema.create_index(
527        key="sparse_metadata",
528        config=SparseVectorIndexConfig(
529            source_key="raw_text",
530            embedding_function=sparse_ef,
531        ),
532    )
533    schema.create_index(key="tag_a", config=StringInvertedIndexConfig())
534    schema.create_index(key="tag_b", config=StringInvertedIndexConfig())
535
536    collection, _ = _create_isolated_collection(client_factories, schema=schema)
537
538    assert collection.schema is not None
539    assert "sparse_metadata" in collection.schema.keys
540    assert "tag_a" in collection.schema.keys
541    assert "tag_b" in collection.schema.keys
542
543    collection.add(
544        ids=["sparse-1"],
545        documents=["source document"],
546        metadatas=[{"raw_text": "oranges", "tag_a": "citrus", "tag_b": "fruit"}],
547    )
548
549    stored = collection.get(ids=["sparse-1"], include=["metadatas"])
550    assert stored["metadatas"] is not None
551    metadata = stored["metadatas"][0]
552    assert metadata is not None
553    assert metadata["tag_a"] == "citrus"
554    assert metadata["tag_b"] == "fruit"
555    assert metadata["raw_text"] == "oranges"
556    assert "sparse_metadata" in metadata
557    sparse_payload = cast(SparseVector, metadata["sparse_metadata"])
558    assert sparse_payload == sparse_ef(["oranges"])[0]
559
560    search_result = collection.search(
561        Search().rank(Knn(key="sparse_metadata", query=cast(Any, sparse_payload)))
562    )
563    assert len(search_result["ids"]) == 1
564    assert "sparse-1" in search_result["ids"][0]
565
566    with pytest.raises(ValueError):
567        collection.schema.create_index(
568            key="another_sparse",
569            config=SparseVectorIndexConfig(source_key="raw_text"),
570        )
571
572    string_filter = collection.get(where={"tag_b": "fruit"})
573    assert set(string_filter["ids"]) == {"sparse-1"}
574
575
576def test_schema_persistence_with_custom_overrides(
577    client_factories: "ClientFactories",
578) -> None:
579    """Custom schema overrides persist across new client instances."""
580    schema = Schema()
581    schema.create_index(key="title", config=StringInvertedIndexConfig())
582    schema.create_index(key="published_year", config=IntInvertedIndexConfig())
583    schema.create_index(key="score", config=FloatInvertedIndexConfig())
584    schema.create_index(key="is_featured", config=BoolInvertedIndexConfig())
585
586    collection, client = _create_isolated_collection(
587        client_factories,
588        schema=schema,
589    )
590
591    collection.add(
592        ids=["persist-1"],
593        documents=["persistent doc"],
594        metadatas=[
595            {
596                "title": "Schema Persistence",
597                "published_year": 2024,
598                "score": 4.5,
599                "is_featured": True,
600            }
601        ],
602    )
603
604    assert collection.schema is not None
605    expected_schema_json = collection.schema.serialize_to_json()
606
607    reloaded_client = client_factories.create_client_from_system()
608    reloaded_collection = reloaded_client.get_collection(name=collection.name)
609    assert reloaded_collection.schema is not None
610    if not is_spann_disabled_mode:
611        assert reloaded_collection.schema.serialize_to_json() == expected_schema_json
612
613    fetched = reloaded_collection.get(where={"title": "Schema Persistence"})
614    assert set(fetched["ids"]) == {"persist-1"}
615
616
617def test_collection_embed_uses_schema_or_collection_embedding_function(
618    client_factories: "ClientFactories",
619) -> None:
620    """_embed should respect schema-provided and direct embedding functions."""
621
622    schema_emb_fn = SimpleEmbeddingFunction(dim=5)
623    schema = Schema().create_index(
624        config=VectorIndexConfig(embedding_function=schema_emb_fn)
625    )
626    schema_collection, _ = _create_isolated_collection(
627        client_factories,
628        schema=schema,
629    )
630
631    schema_embeddings = schema_collection._embed(["schema document"])
632    assert schema_embeddings is not None
633    assert np.allclose(schema_embeddings[0], [0.0, 1.0, 2.0, 3.0, 4.0])
634
635    direct_emb_fn = SimpleEmbeddingFunction(dim=3)
636    direct_collection, _ = _create_isolated_collection(
637        client_factories,
638        embedding_function=direct_emb_fn,
639    )
640
641    direct_embeddings = direct_collection._embed(["direct document"])
642    assert direct_embeddings is not None
643    assert np.allclose(direct_embeddings[0], [0.0, 1.0, 2.0])
644
645
646@pytest.mark.skipif(is_spann_disabled_mode, reason=skip_reason_spann_disabled)
647def test_search_embeds_string_knn_queries(
648    client_factories: "ClientFactories",
649) -> None:
650    """_embed_search_string_queries should embed string KNN queries using collection EF."""
651
652    embedding_fn = RecordingSearchEmbeddingFunction(label="primary")
653    collection, _ = _create_isolated_collection(
654        client_factories, embedding_function=embedding_fn
655    )
656
657    search = Search().rank(Knn(query="hello world"))
658
659    print(collection.schema)
660
661    embedded_search = collection._embed_search_string_queries(search)
662
663    assert embedding_fn.query_inputs == [["hello world"]]
664    assert not embedding_fn.call_inputs
665
666    assert isinstance(search._rank, Knn)
667    assert search._rank.query == "hello world"
668
669    embedded_rank = embedded_search._rank
670    assert isinstance(embedded_rank, Knn)
671    assert embedded_rank.query == [11.0, 12.5]
672
673
674@pytest.mark.skipif(is_spann_disabled_mode, reason=skip_reason_spann_disabled)
675def test_search_embeds_string_knn_queries_with_sparse_embedding_function(
676    client_factories: "ClientFactories",
677) -> None:
678    """_embed_search_string_queries should embed string KNN queries using collection EF."""
679
680    sparse_ef = DeterministicSparseEmbeddingFunction(label="sparse")
681    schema = Schema().create_index(
682        key="sparse_metadata",
683        config=SparseVectorIndexConfig(
684            source_key="raw_text", embedding_function=sparse_ef
685        ),
686    )
687    collection, _ = _create_isolated_collection(client_factories, schema=schema)
688
689    search = Search().rank(Knn(key="sparse_metadata", query="hello world"))
690
691    embedded_search = collection._embed_search_string_queries(search)
692
693    assert isinstance(search._rank, Knn)
694    assert search._rank.key == "sparse_metadata"
695    assert search._rank.query == "hello world"
696
697    embedded_rank = embedded_search._rank
698    assert isinstance(embedded_rank, Knn)
699    assert embedded_rank.key == "sparse_metadata"
700    print(embedded_rank.query)
701    assert embedded_rank.query == SparseVector(indices=[0], values=[11.0])
702
703
704@pytest.mark.skipif(is_spann_disabled_mode, reason=skip_reason_spann_disabled)
705def test_search_embeds_string_queries_in_nested_ranks(
706    client_factories: "ClientFactories",
707) -> None:
708    """String queries in composite rank trees should all be embedded."""
709
710    embedding_fn = RecordingSearchEmbeddingFunction(label="nested")
711    collection, _ = _create_isolated_collection(
712        client_factories, embedding_function=embedding_fn
713    )
714
715    rank_one = (Knn(query="alpha") + Knn(query="beta")).max(Knn(query="gamma"))
716    rank_two = (Knn(query="delta") / 2).abs()
717
718    searches = [Search().rank(rank_one), Search().rank(rank_two)]
719    embedded_searches = [collection._embed_search_string_queries(s) for s in searches]
720
721    expected_queries = [["alpha"], ["beta"], ["gamma"], ["delta"]]
722    assert embedding_fn.query_inputs == expected_queries
723
724    all_queries = []
725    for embedded_search in embedded_searches:
726        all_queries.extend(_collect_knn_queries(embedded_search._rank))
727
728    assert all_queries
729    assert all(not isinstance(query, str) for query in all_queries)
730    assert all(isinstance(query, list) for query in all_queries)
731
732
733def test_schema_delete_index_and_restore(
734    client_factories: "ClientFactories",
735) -> None:
736    """Toggling inverted index enablement reflects in query behavior."""
737    disabled_defaults = Schema().delete_index(config=StringInvertedIndexConfig())
738    collection, client = _create_isolated_collection(
739        client_factories,
740        schema=disabled_defaults,
741    )
742
743    collection.add(
744        ids=["no-index"],
745        documents=["doc"],
746        metadatas=[{"global_field": "value"}],
747    )
748
749    with pytest.raises(Exception):
750        collection.get(where={"global_field": "value"})
751
752    client.delete_collection(name=collection.name)
753
754    disabled_key_schema = (
755        Schema()
756        .create_index(config=StringInvertedIndexConfig())
757        .delete_index(key="category", config=StringInvertedIndexConfig())
758    )
759    recreated = client.get_or_create_collection(
760        name=collection.name, schema=disabled_key_schema
761    )
762    recreated.add(
763        ids=["key-disabled"],
764        documents=["doc"],
765        metadatas=[{"category": "news"}],
766    )
767
768    with pytest.raises(Exception):
769        recreated.get(where={"category": "news"})
770
771    client.delete_collection(name=collection.name)
772
773    restored_schema = Schema().create_index(
774        key="category", config=StringInvertedIndexConfig()
775    )
776    restored = client.get_or_create_collection(
777        name=collection.name, schema=restored_schema
778    )
779    restored.add(
780        ids=["key-enabled"],
781        documents=["doc"],
782        metadatas=[{"category": "news"}],
783    )
784
785    search = restored.get(where={"category": "news"})
786    assert set(search["ids"]) == {"key-enabled"}
787
788
789def test_disabled_metadata_index_filters_raise_invalid_argument_all_modes(
790    client_factories: "ClientFactories",
791) -> None:
792    """Disabled metadata inverted index should block filter-based operations in get, query, and delete for local, single node, and distributed."""
793    schema = Schema().delete_index(
794        key="restricted_tag", config=StringInvertedIndexConfig()
795    )
796    collection, _ = _create_isolated_collection(client_factories, schema=schema)
797
798    collection.add(
799        ids=["restricted-doc"],
800        embeddings=cast(Embeddings, [[0.1, 0.2, 0.3, 0.4]]),
801        metadatas=[{"restricted_tag": "blocked"}],
802        documents=["doc"],
803    )
804
805    assert collection.schema is not None
806    schema_entry = collection.schema.keys["restricted_tag"].string
807    assert schema_entry is not None
808    index_config = schema_entry.string_inverted_index
809    assert index_config is not None
810    assert index_config.enabled is False
811
812    filter_payload: Dict[str, Any] = {"restricted_tag": "blocked"}
813
814    def _expect_disabled_error(operation: Callable[[], Any]) -> None:
815        with pytest.raises(InvalidArgumentError) as exc_info:
816            operation()
817        assert "Cannot filter using metadata key 'restricted_tag'" in str(
818            exc_info.value
819        )
820
821    operations: List[Callable[[], Any]] = [
822        lambda: collection.get(where=filter_payload),
823        lambda: collection.query(
824            query_embeddings=cast(Embeddings, [[0.1, 0.2, 0.3, 0.4]]),
825            n_results=1,
826            where=filter_payload,
827        ),
828        lambda: collection.delete(where=filter_payload),
829    ]
830
831    for operation in operations:
832        _expect_disabled_error(operation)
833
834
835@pytest.mark.skipif(is_spann_disabled_mode, reason=skip_reason_spann_disabled)
836def test_disabled_metadata_index_filters_raise_invalid_argument(
837    client_factories: "ClientFactories",
838) -> None:
839    """Disabled metadata inverted index should block filter-based operations."""
840    schema = Schema().delete_index(
841        key="restricted_tag", config=StringInvertedIndexConfig()
842    )
843    collection, _ = _create_isolated_collection(client_factories, schema=schema)
844
845    collection.add(
846        ids=["restricted-doc"],
847        embeddings=cast(Embeddings, [[0.1, 0.2, 0.3, 0.4]]),
848        metadatas=[{"restricted_tag": "blocked"}],
849        documents=["doc"],
850    )
851
852    assert collection.schema is not None
853    schema_entry = collection.schema.keys["restricted_tag"].string
854    assert schema_entry is not None
855    index_config = schema_entry.string_inverted_index
856    assert index_config is not None
857    assert index_config.enabled is False
858
859    filter_payload: Dict[str, Any] = {"restricted_tag": "blocked"}
860    search_request = Search(where=filter_payload)
861
862    def _expect_disabled_error(operation: Callable[[], Any]) -> None:
863        with pytest.raises(InvalidArgumentError) as exc_info:
864            operation()
865        assert "Cannot filter using metadata key 'restricted_tag'" in str(
866            exc_info.value
867        )
868
869    operations: List[Callable[[], Any]] = [
870        lambda: collection.get(where=filter_payload),
871        lambda: collection.query(
872            query_embeddings=cast(Embeddings, [[0.1, 0.2, 0.3, 0.4]]),
873            n_results=1,
874            where=filter_payload,
875        ),
876        lambda: collection.search(search_request),
877        lambda: collection.delete(where=filter_payload),
878    ]
879
880    for operation in operations:
881        _expect_disabled_error(operation)
882
883
884def test_schema_discovers_new_keys_after_compaction(
885    client_factories: "ClientFactories",
886) -> None:
887    """Compaction promotes unseen metadata keys into discoverable schema entries."""
888    collection, client = _create_isolated_collection(client_factories)
889
890    initial_version = get_collection_version(client, collection.name)
891
892    batch_size = 251
893    ids = [f"discover-add-{i}" for i in range(batch_size)]
894    documents = [f"doc {i}" for i in range(batch_size)]
895    metadatas: List[Mapping[str, Any]] = [
896        {"discover_add": f"topic_{i}"} for i in range(batch_size)
897    ]
898
899    collection.add(ids=ids, documents=documents, metadatas=metadatas)
900
901    if not is_spann_disabled_mode:
902        wait_for_version_increase(client, collection.name, initial_version)
903
904    reloaded = client.get_collection(collection.name)
905    assert reloaded.schema is not None
906    assert "discover_add" in reloaded.schema.keys
907    discover_add_config = reloaded.schema.keys["discover_add"].string
908    assert discover_add_config is not None
909    string_inverted_index = discover_add_config.string_inverted_index
910    assert string_inverted_index is not None
911    assert string_inverted_index.enabled is True
912
913    next_version = get_collection_version(client, collection.name)
914
915    upsert_count = 260
916    upsert_ids = [f"discover-upsert-{i}" for i in range(upsert_count)]
917    upsert_docs = [f"upsert doc {i}" for i in range(upsert_count)]
918    upsert_metadatas: List[Mapping[str, Any]] = [
919        {"discover_upsert": f"topic_{i}"} for i in range(upsert_count)
920    ]
921
922    collection.upsert(
923        ids=upsert_ids,
924        documents=upsert_docs,
925        metadatas=upsert_metadatas,
926    )
927
928    if not is_spann_disabled_mode:
929        wait_for_version_increase(client, collection.name, next_version)
930
931    post_upsert = client.get_collection(collection.name)
932    assert post_upsert.schema is not None
933    assert "discover_upsert" in post_upsert.schema.keys
934    discover_upsert_config = post_upsert.schema.keys["discover_upsert"].string
935    assert discover_upsert_config is not None
936    upsert_inverted_index = discover_upsert_config.string_inverted_index
937    assert upsert_inverted_index is not None
938    assert upsert_inverted_index.enabled is True
939
940    result = collection.get(where={"discover_add": "topic_42"})
941    assert set(result["ids"]) == {"discover-add-42"}
942
943    result_upsert = collection.get(where={"discover_upsert": "topic_42"})
944    assert set(result_upsert["ids"]) == {"discover-upsert-42"}
945
946    reload_client = client_factories.create_client_from_system()
947    persisted = reload_client.get_collection(collection.name)
948    assert persisted.schema is not None
949    assert "discover_add" in persisted.schema.keys
950    assert "discover_upsert" in persisted.schema.keys
951
952
953def test_schema_rejects_conflicting_discoverable_key_types(
954    client_factories: "ClientFactories",
955) -> None:
956    """Conflicting value types should not corrupt discoverable schema entries."""
957    collection, client = _create_isolated_collection(client_factories)
958
959    initial_version = get_collection_version(client, collection.name)
960
961    ids = [f"conflict-{i}" for i in range(251)]
962    metadatas: List[Mapping[str, Any]] = [
963        {"conflict_key": f"value_{i}"} for i in range(251)
964    ]
965    documents = [f"doc {i}" for i in range(251)]
966    collection.add(ids=ids, documents=documents, metadatas=metadatas)
967
968    if not is_spann_disabled_mode:
969        wait_for_version_increase(client, collection.name, initial_version)
970
971    collection.upsert(
972        ids=["conflict-bad"],
973        documents=["bad doc"],
974        metadatas=[{"conflict_key": 100}],
975    )
976
977    collection.update(
978        ids=["conflict-0"],
979        metadatas=[{"conflict_key": 200}],
980    )
981
982    schema = client.get_collection(collection.name).schema
983    assert schema is not None
984    assert "conflict_key" in schema.keys
985    conflict_entry = schema.keys["conflict_key"]
986    if (
987        conflict_entry.string is not None
988        and conflict_entry.string.string_inverted_index is not None
989    ):
990        assert conflict_entry.string.string_inverted_index.enabled is True
991
992    fetch = collection.get(where={"conflict_key": "value_10"})
993    assert set(fetch["ids"]) == {"conflict-10"}
994
995    conflict_bad_meta = collection.get(ids=["conflict-bad"], include=["metadatas"])
996    assert conflict_bad_meta["metadatas"] is not None
997    bad_metadata = conflict_bad_meta["metadatas"][0]
998    assert bad_metadata is not None
999    assert isinstance(bad_metadata["conflict_key"], (int, float))
1000
1001
1002@skip_if_not_cluster()
1003@pytest.mark.skipif(is_spann_disabled_mode, reason=skip_reason_spann_disabled)
1004def test_collection_fork_inherits_and_isolates_schema(
1005    client_factories: "ClientFactories",
1006) -> None:
1007    """Assert forked collections inherit schema and evolve independently of the parent."""
1008    schema = Schema()
1009    schema.create_index(key="shared_key", config=StringInvertedIndexConfig())
1010
1011    collection, client = _create_isolated_collection(
1012        client_factories,
1013        schema=schema,
1014    )
1015
1016    parent_version_before_add = get_collection_version(client, collection.name)
1017
1018    parent_ids = [f"parent-{i}" for i in range(251)]
1019    parent_docs = [f"parent doc {i}" for i in range(251)]
1020    parent_metadatas: List[Mapping[str, Any]] = [
1021        {"shared_key": f"parent_{i}"} for i in range(251)
1022    ]
1023
1024    collection.add(
1025        ids=parent_ids,
1026        documents=parent_docs,
1027        metadatas=parent_metadatas,
1028    )
1029
1030    # Wait for parent to compact before forking. Otherwise, the fork inherits
1031    # uncompacted logs, and compaction of those inherited logs could increment
1032    # the fork's version before the fork's own data is compacted.
1033    wait_for_version_increase(client, collection.name, parent_version_before_add)
1034
1035    assert collection.schema is not None
1036    parent_schema_json = collection.schema.serialize_to_json()
1037
1038    fork_name = f"{collection.name}_fork"
1039    forked = collection.fork(fork_name)
1040
1041    assert forked.schema is not None
1042    assert forked.schema.serialize_to_json() == parent_schema_json
1043
1044    fork_version = get_collection_version(client, forked.name)
1045
1046    fork_ids = [f"fork-{i}" for i in range(251)]
1047    fork_docs = [f"fork doc {i}" for i in range(251)]
1048    fork_metadatas: List[Mapping[str, Any]] = [
1049        {"shared_key": "parent", "child_only": f"value_{i}"} for i in range(251)
1050    ]
1051    forked.upsert(ids=fork_ids, documents=fork_docs, metadatas=fork_metadatas)
1052
1053    wait_for_version_increase(client, forked.name, fork_version)
1054
1055    updated_child = client.get_collection(forked.name)
1056    assert updated_child.schema is not None
1057    assert "child_only" in updated_child.schema.keys
1058    child_only_config = updated_child.schema.keys["child_only"].string
1059    assert child_only_config is not None
1060    child_inverted_index = child_only_config.string_inverted_index
1061    assert child_inverted_index is not None
1062    assert child_inverted_index.enabled is True
1063
1064    reloaded_parent = client.get_collection(collection.name)
1065    assert reloaded_parent.schema is not None
1066    assert "child_only" not in reloaded_parent.schema.keys
1067
1068    parent_results = reloaded_parent.get(where={"shared_key": "parent_10"})
1069    assert set(parent_results["ids"]) == {"parent-10"}
1070
1071    child_results = forked.get(where={"child_only": "value_10"})
1072    assert set(child_results["ids"]) == {"fork-10"}
1073
1074
1075@pytest.mark.skipif(is_spann_disabled_mode, reason=skip_reason_spann_disabled)
1076def test_schema_embedding_configuration_enforced(
1077    client_factories: "ClientFactories",
1078) -> None:
1079    """Schema-provided embedding functions drive both dense and sparse embeddings."""
1080    vector_schema = Schema().create_index(
1081        config=VectorIndexConfig(embedding_function=SimpleEmbeddingFunction(dim=5))
1082    )
1083    vector_collection, _ = _create_isolated_collection(
1084        client_factories,
1085        schema=vector_schema,
1086        embedding_function=SimpleEmbeddingFunction(dim=5),
1087    )
1088
1089    vector_collection.add(
1090        ids=["embed-1"],
1091        documents=["embedding document"],
1092    )
1093
1094    embedded = vector_collection.get(ids=["embed-1"], include=["embeddings"])
1095    assert embedded["embeddings"] is not None
1096    assert np.allclose(embedded["embeddings"][0], [0.0, 1.0, 2.0, 3.0, 4.0])
1097
1098    sparse_ef = DeterministicSparseEmbeddingFunction()
1099    sparse_schema = Schema().create_index(
1100        key="sparse_auto",
1101        config=SparseVectorIndexConfig(
1102            source_key="text_to_embed",
1103            embedding_function=sparse_ef,
1104        ),
1105    )
1106    sparse_collection, _ = _create_isolated_collection(
1107        client_factories,
1108        schema=sparse_schema,
1109    )
1110
1111    sparse_collection.add(
1112        ids=["sparse-text"],
1113        documents=["doc"],
1114        metadatas=[{"text_to_embed": "schema embedding"}],
1115    )
1116    sparse_query = sparse_ef(["schema embedding"])[0]
1117    sparse_result = sparse_collection.get(ids=["sparse-text"], include=["metadatas"])
1118    assert sparse_result["metadatas"] is not None
1119    sparse_meta = sparse_result["metadatas"][0]
1120    assert sparse_meta is not None
1121    assert "sparse_auto" in sparse_meta
1122    sparse_payload = cast(SparseVector, sparse_meta["sparse_auto"])
1123    assert sparse_payload == sparse_query
1124
1125    sparse_search = sparse_collection.search(
1126        Search().rank(Knn(key="sparse_auto", query=cast(Any, sparse_payload)))
1127    )
1128    assert len(sparse_search["ids"]) == 1
1129    assert "sparse-text" in sparse_search["ids"][0]
1130
1131    sparse_collection.add(
1132        ids=["sparse-numeric"],
1133        documents=["doc"],
1134        metadatas=[{"text_to_embed": 5}],
1135    )
1136
1137    numeric_meta = sparse_collection.get(ids=["sparse-numeric"], include=["metadatas"])
1138    assert numeric_meta["metadatas"] is not None
1139    numeric_metadata = numeric_meta["metadatas"][0]
1140    assert numeric_metadata is not None
1141    assert "sparse_auto" not in numeric_metadata
1142
1143
1144def test_schema_precedence_for_overrides_discoverables_and_defaults(
1145    client_factories: "ClientFactories",
1146) -> None:
1147    """Explicit overrides take precedence over disabled defaults and discoverables."""
1148    schema = (
1149        Schema()
1150        .delete_index(config=StringInvertedIndexConfig())
1151        .create_index(key="explicit_key", config=StringInvertedIndexConfig())
1152    )
1153
1154    collection, client = _create_isolated_collection(
1155        client_factories,
1156        schema=schema,
1157    )
1158
1159    ids = [f"precedence-{i}" for i in range(260)]
1160    documents = [f"doc {i}" for i in range(260)]
1161    metadatas: List[Mapping[str, Any]] = [
1162        {"explicit_key": "explicit", "discover_key": f"discover_{i}"}
1163        for i in range(260)
1164    ]
1165
1166    initial_version = get_collection_version(client, collection.name)
1167    collection.add(ids=ids, documents=documents, metadatas=metadatas)
1168
1169    if not is_spann_disabled_mode:
1170        wait_for_version_increase(client, collection.name, initial_version)
1171
1172    schema_state = client.get_collection(collection.name).schema
1173    assert schema_state is not None
1174    assert "explicit_key" in schema_state.keys
1175    explicit_key_string = schema_state.keys["explicit_key"].string
1176    assert explicit_key_string is not None
1177    explicit_inverted_index = explicit_key_string.string_inverted_index
1178    assert explicit_inverted_index is not None
1179    assert explicit_inverted_index.enabled is True
1180
1181    assert "discover_key" in schema_state.keys
1182    discover_key_string = schema_state.keys["discover_key"].string
1183    assert discover_key_string is not None
1184    discover_inverted_index = discover_key_string.string_inverted_index
1185    assert discover_inverted_index is not None
1186    assert discover_inverted_index.enabled is False
1187
1188    explicit_result = collection.get(where={"explicit_key": "explicit"})
1189    assert set(explicit_result["ids"]) == set(ids)
1190
1191    with pytest.raises(Exception):
1192        collection.get(where={"discover_key": "discover_5"})
1193
1194
1195@pytest.mark.skipif(is_spann_disabled_mode, reason=skip_reason_spann_disabled)
1196def test_sparse_auto_embedding_with_document_source_no_metadata(
1197    client_factories: "ClientFactories",
1198) -> None:
1199    """Test sparse embedding auto-generation using #document as source with no metadata."""
1200    sparse_ef = DeterministicSparseEmbeddingFunction(label="doc_no_meta")

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

codekingpro/portable-devtools · Team Ai