codekingpro/portable-devtools
114k
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")
