codekingpro/portable-devtools
114k
1from typing import TypedDict, Dict, Any, Optional, cast, get_args
2import json
3from chromadb.api.types import (
4 Space,
5 CollectionMetadata,
6 UpdateMetadata,
7 EmbeddingFunction,
8)
9from chromadb.utils.embedding_functions import (
10 known_embedding_functions,
11 register_embedding_function,
12)
13from multiprocessing import cpu_count
14import warnings
15
16from chromadb.api.types import Schema
17
18
19class HNSWConfiguration(TypedDict, total=False):
20 space: Space
21 ef_construction: int
22 max_neighbors: int
23 ef_search: int
24 num_threads: int
25 batch_size: int
26 sync_threshold: int
27 resize_factor: float
28
29
30class SpannConfiguration(TypedDict, total=False):
31 search_nprobe: int
32 write_nprobe: int
33 space: Space
34 ef_construction: int
35 ef_search: int
36 max_neighbors: int
37 reassign_neighbor_count: int
38 split_threshold: int
39 merge_threshold: int
40
41
42class CollectionConfiguration(TypedDict, total=True):
43 hnsw: Optional[HNSWConfiguration]
44 spann: Optional[SpannConfiguration]
45 embedding_function: Optional[EmbeddingFunction] # type: ignore
46
47
48def load_collection_configuration_from_json_str(
49 config_json_str: str,
50) -> CollectionConfiguration:
51 config_json_map = json.loads(config_json_str)
52 return load_collection_configuration_from_json(config_json_map)
53
54
55# TODO: make warnings prettier and add link to migration docs
56def load_collection_configuration_from_json(
57 config_json_map: Dict[str, Any]
58) -> CollectionConfiguration:
59 if (
60 config_json_map.get("spann") is not None
61 and config_json_map.get("hnsw") is not None
62 ):
63 raise ValueError("hnsw and spann cannot both be provided")
64
65 hnsw_config = None
66 spann_config = None
67 ef_config = None
68
69 # Process vector index configuration (HNSW or SPANN)
70 if config_json_map.get("hnsw") is not None:
71 hnsw_config = cast(HNSWConfiguration, config_json_map["hnsw"])
72 if config_json_map.get("spann") is not None:
73 spann_config = cast(SpannConfiguration, config_json_map["spann"])
74
75 # Process embedding function configuration
76 if config_json_map.get("embedding_function") is not None:
77 ef_config = config_json_map["embedding_function"]
78 if ef_config["type"] == "legacy":
79 warnings.warn(
80 "legacy embedding function config",
81 DeprecationWarning,
82 stacklevel=2,
83 )
84 ef = None
85 else:
86 try:
87 ef_name = ef_config["name"]
88 except KeyError:
89 raise ValueError(
90 f"Embedding function name not found in config: {ef_config}"
91 )
92 try:
93 ef = known_embedding_functions[ef_name]
94 except KeyError:
95 raise ValueError(
96 f"Embedding function {ef_name} not found. Add @register_embedding_function decorator to the class definition."
97 )
98 try:
99 ef = ef.build_from_config(ef_config["config"]) # type: ignore
100 except Exception as e:
101 raise ValueError(
102 f"Could not build embedding function {ef_config['name']} from config {ef_config['config']}: {e}"
103 )
104 else:
105 ef = None
106
107 return CollectionConfiguration(
108 hnsw=hnsw_config,
109 spann=spann_config,
110 embedding_function=ef, # type: ignore
111 )
112
113
114def collection_configuration_to_json_str(config: CollectionConfiguration) -> str:
115 return json.dumps(collection_configuration_to_json(config))
116
117
118def collection_configuration_to_json(config: CollectionConfiguration) -> Dict[str, Any]:
119 if isinstance(config, dict):
120 hnsw_config = config.get("hnsw")
121 spann_config = config.get("spann")
122 ef = config.get("embedding_function")
123 else:
124 try:
125 hnsw_config = config.get_parameter("hnsw").value
126 except ValueError:
127 hnsw_config = None
128 try:
129 spann_config = config.get_parameter("spann").value
130 except ValueError:
131 spann_config = None
132 try:
133 ef = config.get_parameter("embedding_function").value
134 except ValueError:
135 ef = None
136
137 ef_config: Dict[str, Any] | None = None
138 if hnsw_config is not None:
139 try:
140 hnsw_config = cast(HNSWConfiguration, hnsw_config)
141 except Exception as e:
142 raise ValueError(f"not a valid hnsw config: {e}")
143 if spann_config is not None:
144 try:
145 spann_config = cast(SpannConfiguration, spann_config)
146 except Exception as e:
147 raise ValueError(f"not a valid spann config: {e}")
148
149 if ef is None:
150 ef = None
151 ef_config = {"type": "legacy"}
152
153 if ef is not None:
154 try:
155 if ef.is_legacy():
156 ef_config = {"type": "legacy"}
157 else:
158 ef_config = {
159 "name": ef.name(),
160 "type": "known",
161 "config": ef.get_config(),
162 }
163 register_embedding_function(type(ef)) # type: ignore
164 except Exception as e:
165 warnings.warn(
166 f"legacy embedding function config: {e}",
167 DeprecationWarning,
168 stacklevel=2,
169 )
170 ef = None
171 ef_config = {"type": "legacy"}
172
173 return {
174 "hnsw": hnsw_config,
175 "spann": spann_config,
176 "embedding_function": ef_config,
177 }
178
179
180class CreateHNSWConfiguration(TypedDict, total=False):
181 space: Space
182 ef_construction: int
183 max_neighbors: int
184 ef_search: int
185 num_threads: int
186 batch_size: int
187 sync_threshold: int
188 resize_factor: float
189
190
191def json_to_create_hnsw_configuration(
192 json_map: Dict[str, Any]
193) -> CreateHNSWConfiguration:
194 config: CreateHNSWConfiguration = {}
195 if "space" in json_map:
196 space_value = json_map["space"]
197 if space_value in get_args(Space):
198 config["space"] = space_value
199 else:
200 raise ValueError(f"not a valid space: {space_value}")
201 if "ef_construction" in json_map:
202 config["ef_construction"] = json_map["ef_construction"]
203 if "max_neighbors" in json_map:
204 config["max_neighbors"] = json_map["max_neighbors"]
205 if "ef_search" in json_map:
206 config["ef_search"] = json_map["ef_search"]
207 if "num_threads" in json_map:
208 config["num_threads"] = json_map["num_threads"]
209 if "batch_size" in json_map:
210 config["batch_size"] = json_map["batch_size"]
211 if "sync_threshold" in json_map:
212 config["sync_threshold"] = json_map["sync_threshold"]
213 if "resize_factor" in json_map:
214 config["resize_factor"] = json_map["resize_factor"]
215 return config
216
217
218class CreateSpannConfiguration(TypedDict, total=False):
219 search_nprobe: int
220 write_nprobe: int
221 space: Space
222 ef_construction: int
223 ef_search: int
224 max_neighbors: int
225 reassign_neighbor_count: int
226 split_threshold: int
227 merge_threshold: int
228
229
230def json_to_create_spann_configuration(
231 json_map: Dict[str, Any]
232) -> CreateSpannConfiguration:
233 config: CreateSpannConfiguration = {}
234 if "search_nprobe" in json_map:
235 config["search_nprobe"] = json_map["search_nprobe"]
236 if "write_nprobe" in json_map:
237 config["write_nprobe"] = json_map["write_nprobe"]
238 if "space" in json_map:
239 space_value = json_map["space"]
240 if space_value in get_args(Space):
241 config["space"] = space_value
242 else:
243 raise ValueError(f"not a valid space: {space_value}")
244 if "ef_construction" in json_map:
245 config["ef_construction"] = json_map["ef_construction"]
246 if "ef_search" in json_map:
247 config["ef_search"] = json_map["ef_search"]
248 if "max_neighbors" in json_map:
249 config["max_neighbors"] = json_map["max_neighbors"]
250 return config
251
252
253class CreateCollectionConfiguration(TypedDict, total=False):
254 hnsw: Optional[CreateHNSWConfiguration]
255 spann: Optional[CreateSpannConfiguration]
256 embedding_function: Optional[EmbeddingFunction] # type: ignore
257
258
259def create_collection_configuration_from_legacy_collection_metadata(
260 metadata: CollectionMetadata,
261) -> CreateCollectionConfiguration:
262 """Create a CreateCollectionConfiguration from legacy collection metadata"""
263 return create_collection_configuration_from_legacy_metadata_dict(metadata)
264
265
266def create_collection_configuration_from_legacy_metadata_dict(
267 metadata: Dict[str, Any],
268) -> CreateCollectionConfiguration:
269 """Create a CreateCollectionConfiguration from legacy collection metadata"""
270 old_to_new = {
271 "hnsw:space": "space",
272 "hnsw:construction_ef": "ef_construction",
273 "hnsw:M": "max_neighbors",
274 "hnsw:search_ef": "ef_search",
275 "hnsw:num_threads": "num_threads",
276 "hnsw:batch_size": "batch_size",
277 "hnsw:sync_threshold": "sync_threshold",
278 "hnsw:resize_factor": "resize_factor",
279 }
280 json_map = {}
281 for name, value in metadata.items():
282 if name in old_to_new:
283 json_map[old_to_new[name]] = value
284 hnsw_config = json_to_create_hnsw_configuration(json_map)
285 hnsw_config = populate_create_hnsw_defaults(hnsw_config)
286
287 return CreateCollectionConfiguration(hnsw=hnsw_config)
288
289
290# TODO: make warnings prettier and add link to migration docs
291def load_create_collection_configuration_from_json(
292 json_map: Dict[str, Any]
293) -> CreateCollectionConfiguration:
294 if json_map.get("hnsw") is not None and json_map.get("spann") is not None:
295 raise ValueError("hnsw and spann cannot both be provided")
296
297 result = CreateCollectionConfiguration()
298
299 # Handle vector index configuration
300 if json_map.get("hnsw") is not None:
301 result["hnsw"] = json_to_create_hnsw_configuration(json_map["hnsw"])
302
303 if json_map.get("spann") is not None:
304 result["spann"] = json_to_create_spann_configuration(json_map["spann"])
305
306 # Handle embedding function configuration
307 if json_map.get("embedding_function") is not None:
308 ef_config = json_map["embedding_function"]
309 if ef_config["type"] == "legacy":
310 warnings.warn(
311 "legacy embedding function config",
312 DeprecationWarning,
313 stacklevel=2,
314 )
315 else:
316 ef = known_embedding_functions[ef_config["name"]]
317 result["embedding_function"] = ef.build_from_config(ef_config["config"])
318
319 return result
320
321
322def create_collection_configuration_to_json_str(
323 config: CreateCollectionConfiguration,
324 metadata: Optional[CollectionMetadata] = None,
325) -> str:
326 """Convert a CreateCollection configuration to a JSON-serializable string"""
327 return json.dumps(create_collection_configuration_to_json(config, metadata))
328
329
330# TODO: make warnings prettier and add link to migration docs
331def create_collection_configuration_to_json(
332 config: CreateCollectionConfiguration,
333 metadata: Optional[CollectionMetadata] = None,
334) -> Dict[str, Any]:
335 """Convert a CreateCollection configuration to a JSON-serializable dict"""
336 ef_config: Dict[str, Any] | None = None
337 hnsw_config = config.get("hnsw")
338 spann_config = config.get("spann")
339 if hnsw_config is not None:
340 try:
341 hnsw_config = cast(CreateHNSWConfiguration, hnsw_config)
342 except Exception as e:
343 raise ValueError(f"not a valid hnsw config: {e}")
344 if spann_config is not None:
345 try:
346 spann_config = cast(CreateSpannConfiguration, spann_config)
347 except Exception as e:
348 raise ValueError(f"not a valid spann config: {e}")
349
350 if hnsw_config is not None and spann_config is not None:
351 raise ValueError("hnsw and spann cannot both be provided")
352
353 if config.get("embedding_function") is None:
354 ef = None
355 ef_config = {"type": "legacy"}
356 return {
357 "hnsw": hnsw_config,
358 "spann": spann_config,
359 "embedding_function": ef_config,
360 }
361
362 try:
363 ef = cast(EmbeddingFunction, config.get("embedding_function")) # type: ignore
364 if ef.is_legacy():
365 ef_config = {"type": "legacy"}
366 else:
367 # default space logic: if neither hnsw nor spann config is provided and metadata doesn't have space,
368 # then populate space from ef
369 # otherwise dont use default space from ef
370
371 # then validate the space afterwards based on the supported spaces of the embedding function,
372 # warn if space is not supported
373
374 if hnsw_config is None and spann_config is None:
375 if metadata is None or metadata.get("hnsw:space") is None:
376 # this populates space from ef if not provided in either config
377 hnsw_config = CreateHNSWConfiguration(space=ef.default_space())
378
379 # if hnsw config or spann config exists but space is not provided, populate it from ef
380 if hnsw_config is not None and hnsw_config.get("space") is None:
381 hnsw_config["space"] = ef.default_space()
382 if spann_config is not None and spann_config.get("space") is None:
383 spann_config["space"] = ef.default_space()
384
385 # Validate space compatibility with embedding function
386 if hnsw_config is not None:
387 if hnsw_config.get("space") not in ef.supported_spaces():
388 warnings.warn(
389 f"space {hnsw_config.get('space')} is not supported by {ef.name()}. Supported spaces: {ef.supported_spaces()}",
390 UserWarning,
391 stacklevel=2,
392 )
393 if spann_config is not None:
394 if spann_config.get("space") not in ef.supported_spaces():
395 warnings.warn(
396 f"space {spann_config.get('space')} is not supported by {ef.name()}. Supported spaces: {ef.supported_spaces()}",
397 UserWarning,
398 stacklevel=2,
399 )
400
401 # only validate space from metadata if config is not provided
402 if (
403 hnsw_config is None
404 and spann_config is None
405 and metadata is not None
406 and metadata.get("hnsw:space") is not None
407 ):
408 if metadata.get("hnsw:space") not in ef.supported_spaces():
409 warnings.warn(
410 f"space {metadata.get('hnsw:space')} is not supported by {ef.name()}. Supported spaces: {ef.supported_spaces()}",
411 UserWarning,
412 stacklevel=2,
413 )
414
415 ef_config = {
416 "name": ef.name(),
417 "type": "known",
418 "config": ef.get_config(),
419 }
420 register_embedding_function(type(ef)) # type: ignore
421 except Exception as e:
422 warnings.warn(
423 f"legacy embedding function config: {e}",
424 DeprecationWarning,
425 stacklevel=2,
426 )
427 ef = None
428 ef_config = {"type": "legacy"}
429
430 return {
431 "hnsw": hnsw_config,
432 "spann": spann_config,
433 "embedding_function": ef_config,
434 }
435
436
437def populate_create_hnsw_defaults(
438 config: CreateHNSWConfiguration, ef: Optional[EmbeddingFunction] = None # type: ignore
439) -> CreateHNSWConfiguration:
440 """Populate a CreateHNSW configuration with default values"""
441 if config.get("space") is None:
442 config["space"] = ef.default_space() if ef else "l2"
443 if config.get("ef_construction") is None:
444 config["ef_construction"] = 100
445 if config.get("max_neighbors") is None:
446 config["max_neighbors"] = 16
447 if config.get("ef_search") is None:
448 config["ef_search"] = 100
449 if config.get("num_threads") is None:
450 config["num_threads"] = cpu_count()
451 if config.get("batch_size") is None:
452 config["batch_size"] = 100
453 if config.get("sync_threshold") is None:
454 config["sync_threshold"] = 1000
455 if config.get("resize_factor") is None:
456 config["resize_factor"] = 1.2
457 return config
458
459
460class UpdateHNSWConfiguration(TypedDict, total=False):
461 ef_search: int
462 num_threads: int
463 batch_size: int
464 sync_threshold: int
465 resize_factor: float
466
467
468def json_to_update_hnsw_configuration(
469 json_map: Dict[str, Any]
470) -> UpdateHNSWConfiguration:
471 config: UpdateHNSWConfiguration = {}
472 if "ef_search" in json_map:
473 config["ef_search"] = json_map["ef_search"]
474 if "num_threads" in json_map:
475 config["num_threads"] = json_map["num_threads"]
476 if "batch_size" in json_map:
477 config["batch_size"] = json_map["batch_size"]
478 if "sync_threshold" in json_map:
479 config["sync_threshold"] = json_map["sync_threshold"]
480 if "resize_factor" in json_map:
481 config["resize_factor"] = json_map["resize_factor"]
482 return config
483
484
485class UpdateSpannConfiguration(TypedDict, total=False):
486 search_nprobe: int
487 ef_search: int
488
489
490def json_to_update_spann_configuration(
491 json_map: Dict[str, Any]
492) -> UpdateSpannConfiguration:
493 config: UpdateSpannConfiguration = {}
494 if "search_nprobe" in json_map:
495 config["search_nprobe"] = json_map["search_nprobe"]
496 if "ef_search" in json_map:
497 config["ef_search"] = json_map["ef_search"]
498 return config
499
500
501class UpdateCollectionConfiguration(TypedDict, total=False):
502 hnsw: Optional[UpdateHNSWConfiguration]
503 spann: Optional[UpdateSpannConfiguration]
504 embedding_function: Optional[EmbeddingFunction] # type: ignore
505
506
507def update_collection_configuration_from_legacy_collection_metadata(
508 metadata: CollectionMetadata,
509) -> UpdateCollectionConfiguration:
510 """Create an UpdateCollectionConfiguration from legacy collection metadata"""
511 old_to_new = {
512 "hnsw:search_ef": "ef_search",
513 "hnsw:num_threads": "num_threads",
514 "hnsw:batch_size": "batch_size",
515 "hnsw:sync_threshold": "sync_threshold",
516 "hnsw:resize_factor": "resize_factor",
517 }
518 json_map = {}
519 for name, value in metadata.items():
520 if name in old_to_new:
521 json_map[old_to_new[name]] = value
522 hnsw_config = json_to_update_hnsw_configuration(json_map)
523 return UpdateCollectionConfiguration(hnsw=hnsw_config)
524
525
526def update_collection_configuration_from_legacy_update_metadata(
527 metadata: UpdateMetadata,
528) -> UpdateCollectionConfiguration:
529 """Create an UpdateCollectionConfiguration from legacy update metadata"""
530 old_to_new = {
531 "hnsw:search_ef": "ef_search",
532 "hnsw:num_threads": "num_threads",
533 "hnsw:batch_size": "batch_size",
534 "hnsw:sync_threshold": "sync_threshold",
535 "hnsw:resize_factor": "resize_factor",
536 }
537 json_map = {}
538 for name, value in metadata.items():
539 if name in old_to_new:
540 json_map[old_to_new[name]] = value
541 hnsw_config = json_to_update_hnsw_configuration(json_map)
542 return UpdateCollectionConfiguration(hnsw=hnsw_config)
543
544
545def update_collection_configuration_to_json_str(
546 config: UpdateCollectionConfiguration,
547) -> str:
548 """Convert an UpdateCollectionConfiguration to a JSON-serializable string"""
549 json_dict = update_collection_configuration_to_json(config)
550 return json.dumps(json_dict)
551
552
553def update_collection_configuration_to_json(
554 config: UpdateCollectionConfiguration,
555) -> Dict[str, Any]:
556 """Convert an UpdateCollectionConfiguration to a JSON-serializable dict"""
557 hnsw_config = config.get("hnsw")
558 spann_config = config.get("spann")
559 ef = config.get("embedding_function")
560 if hnsw_config is None and spann_config is None and ef is None:
561 return {}
562
563 if hnsw_config is not None:
564 try:
565 hnsw_config = cast(UpdateHNSWConfiguration, hnsw_config)
566 except Exception as e:
567 raise ValueError(f"not a valid hnsw config: {e}")
568
569 if spann_config is not None:
570 try:
571 spann_config = cast(UpdateSpannConfiguration, spann_config)
572 except Exception as e:
573 raise ValueError(f"not a valid spann config: {e}")
574
575 ef_config: Dict[str, Any] | None = None
576 if ef is not None:
577 if ef.is_legacy():
578 ef_config = {"type": "legacy"}
579 else:
580 ef.validate_config(ef.get_config())
581 ef_config = {
582 "name": ef.name(),
583 "type": "known",
584 "config": ef.get_config(),
585 }
586 register_embedding_function(type(ef)) # type: ignore
587 else:
588 ef_config = None
589
590 return {
591 "hnsw": hnsw_config,
592 "spann": spann_config,
593 "embedding_function": ef_config,
594 }
595
596
597def load_update_collection_configuration_from_json_str(
598 json_str: str,
599) -> UpdateCollectionConfiguration:
600 json_map = json.loads(json_str)
601 return load_update_collection_configuration_from_json(json_map)
602
603
604# TODO: make warnings prettier and add link to migration docs
605def load_update_collection_configuration_from_json(
606 json_map: Dict[str, Any]
607) -> UpdateCollectionConfiguration:
608 """Convert a JSON dict to an UpdateCollectionConfiguration"""
609 if json_map.get("hnsw") is not None and json_map.get("spann") is not None:
610 raise ValueError("hnsw and spann cannot both be provided")
611
612 result = UpdateCollectionConfiguration()
613
614 # Handle vector index configurations
615 if json_map.get("hnsw") is not None:
616 result["hnsw"] = json_to_update_hnsw_configuration(json_map["hnsw"])
617
618 if json_map.get("spann") is not None:
619 result["spann"] = json_to_update_spann_configuration(json_map["spann"])
620
621 # Handle embedding function
622 if json_map.get("embedding_function") is not None:
623 if json_map["embedding_function"]["type"] == "legacy":
624 warnings.warn(
625 "legacy embedding function config",
626 DeprecationWarning,
627 stacklevel=2,
628 )
629 else:
630 ef = known_embedding_functions[json_map["embedding_function"]["name"]]
631 result["embedding_function"] = ef.build_from_config(
632 json_map["embedding_function"]["config"]
633 )
634
635 return result
636
637
638def overwrite_hnsw_configuration(
639 existing_hnsw_config: HNSWConfiguration, update_hnsw_config: UpdateHNSWConfiguration
640) -> HNSWConfiguration:
641 """Overwrite a HNSWConfiguration with a new configuration"""
642 # Create a copy of the existing config and update with new values
643 result = dict(existing_hnsw_config)
644 update_fields = [
645 "ef_search",
646 "num_threads",
647 "batch_size",
648 "sync_threshold",
649 "resize_factor",
650 ]
651
652 for field in update_fields:
653 if field in update_hnsw_config:
654 result[field] = update_hnsw_config[field] # type: ignore
655
656 return cast(HNSWConfiguration, result)
657
658
659def overwrite_spann_configuration(
660 existing_spann_config: SpannConfiguration,
661 update_spann_config: UpdateSpannConfiguration,
662) -> SpannConfiguration:
663 """Overwrite a SpannConfiguration with a new configuration"""
664 result = dict(existing_spann_config)
665 update_fields = [
666 "search_nprobe",
667 "ef_search",
668 ]
669
670 for field in update_fields:
671 if field in update_spann_config:
672 result[field] = update_spann_config[field] # type: ignore
673
674 return cast(SpannConfiguration, result)
675
676
677# TODO: make warnings prettier and add link to migration docs
678def overwrite_embedding_function(
679 existing_embedding_function: EmbeddingFunction, # type: ignore
680 update_embedding_function: EmbeddingFunction, # type: ignore
681) -> EmbeddingFunction: # type: ignore
682 """Overwrite an EmbeddingFunction with a new configuration"""
683 # Check for legacy embedding functions
684 if existing_embedding_function.is_legacy() or update_embedding_function.is_legacy():
685 warnings.warn(
686 "cannot update legacy embedding function config",
687 DeprecationWarning,
688 stacklevel=2,
689 )
690 return existing_embedding_function
691
692 # Validate function compatibility
693 if existing_embedding_function.name() != update_embedding_function.name():
694 raise ValueError(
695 f"Cannot update embedding function: incompatible types "
696 f"({existing_embedding_function.name()} vs {update_embedding_function.name()})"
697 )
698
699 # Validate and apply the configuration update
700 update_embedding_function.validate_config_update(
701 existing_embedding_function.get_config(), update_embedding_function.get_config()
702 )
703 return update_embedding_function
704
705
706def overwrite_collection_configuration(
707 existing_config: CollectionConfiguration,
708 update_config: UpdateCollectionConfiguration,
709) -> CollectionConfiguration:
710 """Overwrite a CollectionConfiguration with a new configuration"""
711 update_spann = update_config.get("spann")
712 update_hnsw = update_config.get("hnsw")
713 if update_spann is not None and update_hnsw is not None:
714 raise ValueError("hnsw and spann cannot both be provided")
715
716 # Handle HNSW configuration update
717
718 updated_hnsw_config = existing_config.get("hnsw")
719 if updated_hnsw_config is not None and update_hnsw is not None:
720 updated_hnsw_config = overwrite_hnsw_configuration(
721 updated_hnsw_config, update_hnsw
722 )
723
724 # Handle SPANN configuration update
725 updated_spann_config = existing_config.get("spann")
726 if updated_spann_config is not None and update_spann is not None:
727 updated_spann_config = overwrite_spann_configuration(
728 updated_spann_config, update_spann
729 )
730
731 # Handle embedding function update
732 updated_embedding_function = existing_config.get("embedding_function")
733 update_ef = update_config.get("embedding_function")
734 if update_ef is not None:
735 if updated_embedding_function is not None:
736 updated_embedding_function = overwrite_embedding_function(
737 updated_embedding_function, update_ef
738 )
739 else:
740 updated_embedding_function = update_ef
741
742 return CollectionConfiguration(
743 hnsw=updated_hnsw_config,
744 spann=updated_spann_config,
745 embedding_function=updated_embedding_function,
746 )
747
748
749def validate_embedding_function_conflict_on_create(
750 embedding_function: Optional[EmbeddingFunction], # type: ignore
751 configuration_ef: Optional[EmbeddingFunction], # type: ignore
752) -> None:
753 """
754 Validates that there are no conflicting embedding functions between function parameter
755 and collection configuration.
756
757 Args:
758 embedding_function: The embedding function provided as a parameter
759 configuration_ef: The embedding function from collection configuration
760
761 Returns:
762 bool: True if there is a conflict, False otherwise
763 """
764 # If ef provided in function params and collection config, check if they are the same
765 # If not, there's a conflict
766 # ef is by default "default" if not provided, so ignore that case.
767 if embedding_function is not None and configuration_ef is not None:
768 if (
769 embedding_function.name() != "default"
770 and embedding_function.name() != configuration_ef.name()
771 ):
772 raise ValueError(
773 f"Multiple embedding functions provided. Please provide only one. Embedding function conflict: {embedding_function.name()} vs {configuration_ef.name()}"
774 )
775 return None
776
777
778# The reason to use the config on get, rather than build the ef is because
779# if there is an issue with deserializing the config, an error shouldn't be raised
780# at get time. CollectionCommon.py will raise an error at _embed time if there is an issue deserializing.
781def validate_embedding_function_conflict_on_get(
782 embedding_function: Optional[EmbeddingFunction], # type: ignore
783 persisted_ef_config: Optional[Dict[str, Any]],
784) -> None:
785 """
786 Validates that there are no conflicting embedding functions between function parameter
787 and collection configuration.
788 """
789 if persisted_ef_config is not None and embedding_function is not None:
790 if (
791 embedding_function.name() != "default"
792 and persisted_ef_config.get("name") is not None
793 and persisted_ef_config.get("name") != embedding_function.name()
794 ):
795 raise ValueError(
796 f"An embedding function already exists in the collection configuration, and a new one is provided. If this is intentional, please embed documents separately. Embedding function conflict: new: {embedding_function.name()} vs persisted: {persisted_ef_config.get('name')}"
797 )
798 return None
799
800
801def update_schema_from_collection_configuration(
802 schema: "Schema", configuration: "UpdateCollectionConfiguration"
803) -> "Schema":
804 """
805 Updates a schema with configuration changes.
806 Only updates fields that are present in the configuration update.
807
808 Args:
809 schema: The existing Schema object
810 configuration: The configuration updates to apply
811
812 Returns:
813 Updated Schema object
814 """
815
816 # Get the vector index from defaults and #embedding key
817 if (
818 schema.defaults.float_list is None
819 or schema.defaults.float_list.vector_index is None
820 ):
821 raise ValueError("Schema is missing defaults.float_list.vector_index")
822
823 embedding_key = "#embedding"
824 if embedding_key not in schema.keys:
825 raise ValueError(f"Schema is missing keys[{embedding_key}]")
826
827 embedding_value_types = schema.keys[embedding_key]
828 if (
829 embedding_value_types.float_list is None
830 or embedding_value_types.float_list.vector_index is None
831 ):
832 raise ValueError(
833 f"Schema is missing keys[{embedding_key}].float_list.vector_index"
834 )
835
836 # Update vector index config in both locations
837 for vector_index in [
838 schema.defaults.float_list.vector_index,
839 embedding_value_types.float_list.vector_index,
840 ]:
841 if "hnsw" in configuration and configuration["hnsw"] is not None:
842 # Update HNSW config
843 if vector_index.config.hnsw is None:
844 raise ValueError("Trying to update HNSW config but schema has SPANN")
845
846 hnsw_config = vector_index.config.hnsw
847 update_hnsw = configuration["hnsw"]
848
849 # Only update fields that are present in the update
850 if "ef_search" in update_hnsw:
851 hnsw_config.ef_search = update_hnsw["ef_search"]
852 if "num_threads" in update_hnsw:
853 hnsw_config.num_threads = update_hnsw["num_threads"]
854 if "batch_size" in update_hnsw:
855 hnsw_config.batch_size = update_hnsw["batch_size"]
856 if "sync_threshold" in update_hnsw:
857 hnsw_config.sync_threshold = update_hnsw["sync_threshold"]
858 if "resize_factor" in update_hnsw:
859 hnsw_config.resize_factor = update_hnsw["resize_factor"]
860
861 elif "spann" in configuration and configuration["spann"] is not None:
862 # Update SPANN config
863 if vector_index.config.spann is None:
864 raise ValueError("Trying to update SPANN config but schema has HNSW")
865
866 spann_config = vector_index.config.spann
867 update_spann = configuration["spann"]
868
869 # Only update fields that are present in the update
870 if "search_nprobe" in update_spann:
871 spann_config.search_nprobe = update_spann["search_nprobe"]
872 if "ef_search" in update_spann:
873 spann_config.ef_search = update_spann["ef_search"]
874
875 # Update embedding function if present
876 if (
877 "embedding_function" in configuration
878 and configuration["embedding_function"] is not None
879 ):
880 vector_index.config.embedding_function = configuration["embedding_function"]
881
882 return schema
883 