Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
collection_configuration.py883 linesDownload Raw Back to api
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 
codekingpro/portable-devtools · Team Ai