Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
CollectionCommon.py1040 linesDownload Raw Back to models
1import functools
2from typing import (
3    TYPE_CHECKING,
4    Callable,
5    Dict,
6    Generic,
7    Optional,
8    Any,
9    Set,
10    TypeVar,
11    Union,
12    cast,
13    List,
14)
15from chromadb.types import Metadata
16import numpy as np
17from uuid import UUID
18
19from chromadb.api.types import (
20    URI,
21    Schema,
22    SparseVectorIndexConfig,
23    URIs,
24    AddRequest,
25    BaseRecordSet,
26    CollectionMetadata,
27    DataLoader,
28    DeleteRequest,
29    Embedding,
30    Embeddings,
31    FilterSet,
32    GetRequest,
33    PyEmbedding,
34    Embeddable,
35    GetResult,
36    Include,
37    Loadable,
38    Document,
39    Image,
40    QueryRequest,
41    QueryResult,
42    IDs,
43    EmbeddingFunction,
44    SparseEmbeddingFunction,
45    ID,
46    OneOrMany,
47    UpdateRequest,
48    UpsertRequest,
49    get_default_embeddable_record_set_fields,
50    maybe_cast_one_to_many,
51    normalize_base_record_set,
52    normalize_insert_record_set,
53    validate_base_record_set,
54    validate_ids,
55    validate_include,
56    validate_insert_record_set,
57    validate_metadata,
58    validate_metadatas,
59    validate_embedding_function,
60    validate_sparse_embedding_function,
61    validate_n_results,
62    validate_record_set_contains_any,
63    validate_record_set_for_embedding,
64    validate_filter_set,
65    DefaultEmbeddingFunction,
66    EMBEDDING_KEY,
67    DOCUMENT_KEY,
68)
69from chromadb.api.collection_configuration import (
70    UpdateCollectionConfiguration,
71    overwrite_collection_configuration,
72    load_collection_configuration_from_json,
73    CollectionConfiguration,
74)
75
76# TODO: We should rename the types in chromadb.types to be Models where
77# appropriate. This will help to distinguish between manipulation objects
78# which are essentially API views. And the actual data models which are
79# stored / retrieved / transmitted.
80from chromadb.types import Collection as CollectionModel, Where, WhereDocument
81import logging
82
83logger = logging.getLogger(__name__)
84
85if TYPE_CHECKING:
86    from chromadb.api import ServerAPI, AsyncServerAPI
87
88ClientT = TypeVar("ClientT", "ServerAPI", "AsyncServerAPI")
89
90T = TypeVar("T")
91
92
93def validation_context(name: str) -> Callable[[Callable[..., T]], Callable[..., T]]:
94    """A decorator that wraps a method with a try-except block that catches
95    exceptions and adds the method name to the error message. This allows us to
96    provide more context when an error occurs, without rewriting validators.
97    """
98
99    def decorator(func: Callable[..., T]) -> Callable[..., T]:
100        @functools.wraps(func)
101        def wrapper(self: Any, *args: Any, **kwargs: Any) -> T:
102            try:
103                return func(self, *args, **kwargs)
104            except Exception as e:
105                msg = f"{str(e)} in {name}."
106                # add the rest of the args to the error message if they exist
107                e.args = (msg,) + e.args[1:] if e.args else ()
108                # raise the same error that was caught with the modified message
109                raise
110
111        return wrapper
112
113    return decorator
114
115
116class CollectionCommon(Generic[ClientT]):
117    _model: CollectionModel
118    _client: ClientT
119    _embedding_function: Optional[EmbeddingFunction[Embeddable]]
120    _data_loader: Optional[DataLoader[Loadable]]
121
122    def __init__(
123        self,
124        client: ClientT,
125        model: CollectionModel,
126        embedding_function: Optional[
127            EmbeddingFunction[Embeddable]
128        ] = DefaultEmbeddingFunction(),  # type: ignore
129        data_loader: Optional[DataLoader[Loadable]] = None,
130    ):
131        """Initializes a new instance of the Collection class."""
132
133        self._client = client
134        self._model = model
135
136        # Check to make sure the embedding function has the right signature, as defined by the EmbeddingFunction protocol
137        if embedding_function is not None:
138            validate_embedding_function(embedding_function)
139
140        self._embedding_function = embedding_function
141        self._data_loader = data_loader
142
143    # Expose the model properties as read-only properties on the Collection class
144
145    @property
146    def id(self) -> UUID:
147        return self._model.id
148
149    @property
150    def name(self) -> str:
151        return self._model.name
152
153    @property
154    def configuration(self) -> CollectionConfiguration:
155        return load_collection_configuration_from_json(self._model.configuration_json)
156
157    @property
158    def configuration_json(self) -> Dict[str, Any]:
159        return self._model.configuration_json
160
161    @property
162    def schema(self) -> Optional[Schema]:
163        return Schema.deserialize_from_json(
164            self._model.serialized_schema if self._model.serialized_schema else {}
165        )
166
167    @property
168    def metadata(self) -> CollectionMetadata:
169        return cast(CollectionMetadata, self._model.metadata)
170
171    @property
172    def tenant(self) -> str:
173        return self._model.tenant
174
175    @property
176    def database(self) -> str:
177        return self._model.database
178
179    def __eq__(self, other: object) -> bool:
180        if not isinstance(other, CollectionCommon):
181            return False
182        id_match = self.id == other.id
183        name_match = self.name == other.name
184        configuration_match = self.configuration_json == other.configuration_json
185        schema_match = self.schema == other.schema
186        metadata_match = self.metadata == other.metadata
187        tenant_match = self.tenant == other.tenant
188        database_match = self.database == other.database
189        embedding_function_match = self._embedding_function == other._embedding_function
190        data_loader_match = self._data_loader == other._data_loader
191        return (
192            id_match
193            and name_match
194            and configuration_match
195            and schema_match
196            and metadata_match
197            and tenant_match
198            and database_match
199            and embedding_function_match
200            and data_loader_match
201        )
202
203    def __repr__(self) -> str:
204        return f"Collection(name={self.name})"
205
206    def get_model(self) -> CollectionModel:
207        return self._model
208
209    @validation_context("add")
210    def _validate_and_prepare_add_request(
211        self,
212        ids: OneOrMany[ID],
213        embeddings: Optional[
214            Union[
215                OneOrMany[Embedding],
216                OneOrMany[PyEmbedding],
217            ]
218        ],
219        metadatas: Optional[OneOrMany[Metadata]],
220        documents: Optional[OneOrMany[Document]],
221        images: Optional[OneOrMany[Image]],
222        uris: Optional[OneOrMany[URI]],
223    ) -> AddRequest:
224        # Unpack
225        add_records = normalize_insert_record_set(
226            ids=ids,
227            embeddings=embeddings,
228            metadatas=metadatas,
229            documents=documents,
230            images=images,
231            uris=uris,
232        )
233
234        # Validate
235        validate_insert_record_set(record_set=add_records)
236        validate_record_set_contains_any(record_set=add_records, contains_any={"ids"})
237
238        # Prepare
239        if add_records["embeddings"] is None:
240            validate_record_set_for_embedding(record_set=add_records)
241            add_embeddings = self._embed_record_set(record_set=add_records)
242        else:
243            add_embeddings = add_records["embeddings"]
244
245        add_metadatas = self._apply_sparse_embeddings_to_metadatas(
246            add_records["metadatas"], add_records["documents"]
247        )
248
249        return AddRequest(
250            ids=add_records["ids"],
251            embeddings=add_embeddings,
252            metadatas=add_metadatas,
253            documents=add_records["documents"],
254            uris=add_records["uris"],
255        )
256
257    @validation_context("get")
258    def _validate_and_prepare_get_request(
259        self,
260        ids: Optional[OneOrMany[ID]],
261        where: Optional[Where],
262        where_document: Optional[WhereDocument],
263        include: Include,
264    ) -> GetRequest:
265        # Unpack
266        unpacked_ids: Optional[IDs] = maybe_cast_one_to_many(target=ids)
267        filters = FilterSet(where=where, where_document=where_document)
268
269        # Validate
270        if unpacked_ids is not None:
271            validate_ids(ids=unpacked_ids)
272
273        validate_filter_set(filter_set=filters)
274        validate_include(include=include, dissalowed=["distances"])
275
276        if "data" in include and self._data_loader is None:
277            raise ValueError(
278                "You must set a data loader on the collection if loading from URIs."
279            )
280
281        # Prepare
282        request_include = include
283        # We need to include uris in the result from the API to load datas
284        if "data" in include and "uris" not in include:
285            request_include.append("uris")
286
287        return GetRequest(
288            ids=unpacked_ids,
289            where=filters["where"],
290            where_document=filters["where_document"],
291            include=request_include,
292        )
293
294    @validation_context("query")
295    def _validate_and_prepare_query_request(
296        self,
297        query_embeddings: Optional[
298            Union[
299                OneOrMany[Embedding],
300                OneOrMany[PyEmbedding],
301            ]
302        ],
303        query_texts: Optional[OneOrMany[Document]],
304        query_images: Optional[OneOrMany[Image]],
305        query_uris: Optional[OneOrMany[URI]],
306        ids: Optional[OneOrMany[ID]],
307        n_results: int,
308        where: Optional[Where],
309        where_document: Optional[WhereDocument],
310        include: Include,
311    ) -> QueryRequest:
312        # Unpack
313        query_records = normalize_base_record_set(
314            embeddings=query_embeddings,
315            documents=query_texts,
316            images=query_images,
317            uris=query_uris,
318        )
319
320        filter_ids = maybe_cast_one_to_many(ids)
321
322        filters = FilterSet(
323            where=where,
324            where_document=where_document,
325        )
326
327        # Validate
328        validate_base_record_set(record_set=query_records)
329        validate_filter_set(filter_set=filters)
330        validate_include(include=include)
331        validate_n_results(n_results=n_results)
332
333        # Prepare
334        if query_records["embeddings"] is None:
335            validate_record_set_for_embedding(record_set=query_records)
336            request_embeddings = self._embed_record_set(
337                record_set=query_records, is_query=True
338            )
339        else:
340            request_embeddings = query_records["embeddings"]
341
342        request_where = filters["where"]
343        request_where_document = filters["where_document"]
344
345        # We need to manually include uris in the result from the API to load datas
346        request_include = include
347        if "data" in request_include and "uris" not in request_include:
348            request_include.append("uris")
349
350        return QueryRequest(
351            embeddings=request_embeddings,
352            ids=filter_ids,
353            where=request_where,
354            where_document=request_where_document,
355            include=request_include,
356            n_results=n_results,
357        )
358
359    @validation_context("update")
360    def _validate_and_prepare_update_request(
361        self,
362        ids: OneOrMany[ID],
363        embeddings: Optional[
364            Union[
365                OneOrMany[Embedding],
366                OneOrMany[PyEmbedding],
367            ]
368        ],
369        metadatas: Optional[OneOrMany[Metadata]],
370        documents: Optional[OneOrMany[Document]],
371        images: Optional[OneOrMany[Image]],
372        uris: Optional[OneOrMany[URI]],
373    ) -> UpdateRequest:
374        # Unpack
375        update_records = normalize_insert_record_set(
376            ids=ids,
377            embeddings=embeddings,
378            metadatas=metadatas,
379            documents=documents,
380            images=images,
381            uris=uris,
382        )
383
384        # Validate
385        validate_insert_record_set(record_set=update_records)
386
387        # Prepare
388        if update_records["embeddings"] is None:
389            # TODO: Handle URI updates.
390            if (
391                update_records["documents"] is not None
392                or update_records["images"] is not None
393            ):
394                validate_record_set_for_embedding(
395                    update_records, embeddable_fields={"documents", "images"}
396                )
397                update_embeddings = self._embed_record_set(record_set=update_records)
398            else:
399                update_embeddings = None
400        else:
401            update_embeddings = update_records["embeddings"]
402
403        update_metadatas = self._apply_sparse_embeddings_to_metadatas(
404            update_records["metadatas"], update_records["documents"]
405        )
406
407        return UpdateRequest(
408            ids=update_records["ids"],
409            embeddings=update_embeddings,
410            metadatas=update_metadatas,
411            documents=update_records["documents"],
412            uris=update_records["uris"],
413        )
414
415    @validation_context("upsert")
416    def _validate_and_prepare_upsert_request(
417        self,
418        ids: OneOrMany[ID],
419        embeddings: Optional[
420            Union[
421                OneOrMany[Embedding],
422                OneOrMany[PyEmbedding],
423            ]
424        ] = None,
425        metadatas: Optional[OneOrMany[Metadata]] = None,
426        documents: Optional[OneOrMany[Document]] = None,
427        images: Optional[OneOrMany[Image]] = None,
428        uris: Optional[OneOrMany[URI]] = None,
429    ) -> UpsertRequest:
430        # Unpack
431        upsert_records = normalize_insert_record_set(
432            ids=ids,
433            embeddings=embeddings,
434            metadatas=metadatas,
435            documents=documents,
436            images=images,
437            uris=uris,
438        )
439
440        # Validate
441        validate_insert_record_set(record_set=upsert_records)
442
443        # Prepare
444        if upsert_records["embeddings"] is None:
445            validate_record_set_for_embedding(
446                record_set=upsert_records, embeddable_fields={"documents", "images"}
447            )
448            upsert_embeddings = self._embed_record_set(record_set=upsert_records)
449        else:
450            upsert_embeddings = upsert_records["embeddings"]
451
452        upsert_metadatas = self._apply_sparse_embeddings_to_metadatas(
453            upsert_records["metadatas"], upsert_records["documents"]
454        )
455
456        return UpsertRequest(
457            ids=upsert_records["ids"],
458            metadatas=upsert_metadatas,
459            embeddings=upsert_embeddings,
460            documents=upsert_records["documents"],
461            uris=upsert_records["uris"],
462        )
463
464    @validation_context("delete")
465    def _validate_and_prepare_delete_request(
466        self,
467        ids: Optional[IDs],
468        where: Optional[Where],
469        where_document: Optional[WhereDocument],
470        limit: Optional[int] = None,
471    ) -> DeleteRequest:
472        if ids is None and where is None and where_document is None:
473            raise ValueError(
474                "At least one of ids, where, or where_document must be provided"
475            )
476
477        if limit is not None:
478            if not isinstance(limit, int) or isinstance(limit, bool):
479                raise TypeError("limit must be a non-negative integer")
480            if limit < 0:
481                raise ValueError("limit must be a non-negative integer")
482
483        if limit is not None and where is None and where_document is None:
484            raise ValueError(
485                "limit can only be specified when a where or where_document clause is provided"
486            )
487
488        # Unpack
489        if ids is not None:
490            request_ids = cast(IDs, maybe_cast_one_to_many(ids))
491        else:
492            request_ids = None
493        filters = FilterSet(where=where, where_document=where_document)
494
495        # Validate
496        if request_ids is not None:
497            validate_ids(ids=request_ids)
498        validate_filter_set(filter_set=filters)
499
500        return DeleteRequest(
501            ids=request_ids, where=where, where_document=where_document, limit=limit
502        )
503
504    def _transform_peek_response(self, response: GetResult) -> GetResult:
505        if response["embeddings"] is not None:
506            response["embeddings"] = np.array(response["embeddings"])
507
508        return response
509
510    def _transform_get_response(
511        self, response: GetResult, include: Include
512    ) -> GetResult:
513        if (
514            "data" in include
515            and self._data_loader is not None
516            and response["uris"] is not None
517        ):
518            response["data"] = self._data_loader(response["uris"])
519
520        if "embeddings" in include:
521            response["embeddings"] = np.array(response["embeddings"])
522
523        # Remove URIs from the result if they weren't requested
524        if "uris" not in include:
525            response["uris"] = None
526
527        return response
528
529    def _transform_query_response(
530        self, response: QueryResult, include: Include
531    ) -> QueryResult:
532        if (
533            "data" in include
534            and self._data_loader is not None
535            and response["uris"] is not None
536        ):
537            response["data"] = [self._data_loader(uris) for uris in response["uris"]]
538
539        if "embeddings" in include and response["embeddings"] is not None:
540            response["embeddings"] = [
541                np.array(embedding) for embedding in response["embeddings"]
542            ]
543
544        # Remove URIs from the result if they weren't requested
545        if "uris" not in include:
546            response["uris"] = None
547
548        return response
549
550    def _validate_modify_request(self, metadata: Optional[CollectionMetadata]) -> None:
551        if metadata is not None:
552            validate_metadata(metadata)
553            if "hnsw:space" in metadata:
554                raise ValueError(
555                    "Changing the distance function of a collection once it is created is not supported currently."
556                )
557
558    def _update_model_after_modify_success(
559        self,
560        name: Optional[str],
561        metadata: Optional[CollectionMetadata],
562        configuration: Optional[UpdateCollectionConfiguration],
563    ) -> None:
564        if name:
565            self._model["name"] = name
566        if metadata:
567            self._model["metadata"] = metadata
568        if configuration:
569            self._model.set_configuration(
570                overwrite_collection_configuration(
571                    self._model.get_configuration(), configuration
572                )
573            )
574
575            # If schema exists, also update it with the configuration changes
576            if self.schema:
577                from chromadb.api.collection_configuration import (
578                    update_schema_from_collection_configuration,
579                )
580
581                updated_schema = update_schema_from_collection_configuration(
582                    self.schema, configuration
583                )
584                self._model["serialized_schema"] = updated_schema.serialize_to_json()
585
586    def _get_sparse_embedding_targets(self) -> Dict[str, "SparseVectorIndexConfig"]:
587        schema = self.schema
588        if schema is None:
589            return {}
590
591        targets: Dict[str, "SparseVectorIndexConfig"] = {}
592        for key, value_types in schema.keys.items():
593            if value_types.sparse_vector is None:
594                continue
595            sparse_index = value_types.sparse_vector.sparse_vector_index
596            if sparse_index is None or not sparse_index.enabled:
597                continue
598            config = sparse_index.config
599            if config.embedding_function is None or config.source_key is None:
600                continue
601            targets[key] = config
602
603        return targets
604
605    def _apply_sparse_embeddings_to_metadatas(
606        self,
607        metadatas: Optional[List[Metadata]],
608        documents: Optional[List[Document]] = None,
609    ) -> Optional[List[Metadata]]:
610        sparse_targets = self._get_sparse_embedding_targets()
611        if not sparse_targets:
612            return metadatas
613
614        # If no metadatas provided, create empty dicts based on documents length
615        if metadatas is None:
616            if documents is None:
617                return None
618            metadatas = [{} for _ in range(len(documents))]
619
620        # Create copies, converting None to empty dict
621        updated_metadatas: List[Dict[str, Any]] = [
622            dict(metadata) if metadata is not None else {} for metadata in metadatas
623        ]
624
625        documents_list = list(documents) if documents is not None else None
626
627        for target_key, config in sparse_targets.items():
628            source_key = config.source_key
629            embedding_func = config.embedding_function
630            if source_key is None or embedding_func is None:
631                continue
632
633            if not isinstance(embedding_func, SparseEmbeddingFunction):
634                embedding_func = cast(SparseEmbeddingFunction[Any], embedding_func)
635            validate_sparse_embedding_function(embedding_func)
636
637            # Initialize collection lists for batch processing
638            inputs: List[str] = []
639            positions: List[int] = []
640
641            # Handle special case: source_key is "#document"
642            if source_key == DOCUMENT_KEY:
643                if documents_list is None:
644                    continue
645
646                # Collect documents that need embedding
647                for idx, metadata in enumerate(updated_metadatas):
648                    # Skip if target already exists in metadata
649                    if target_key in metadata:
650                        continue
651
652                    # Get document at this position
653                    if idx < len(documents_list):
654                        doc = documents_list[idx]
655                        if isinstance(doc, str):
656                            inputs.append(doc)
657                            positions.append(idx)
658
659                # Generate embeddings for all collected documents
660                if len(inputs) == 0:
661                    continue
662
663                sparse_embeddings = self._sparse_embed(
664                    input=inputs,
665                    sparse_embedding_function=embedding_func,
666                )
667
668                if len(sparse_embeddings) != len(positions):
669                    raise ValueError(
670                        "Sparse embedding function returned unexpected number of embeddings."
671                    )
672
673                for position, embedding in zip(positions, sparse_embeddings):
674                    updated_metadatas[position][target_key] = embedding
675
676                continue  # Skip the metadata-based logic below
677
678            # Handle normal case: source_key is a metadata field
679            for idx, metadata in enumerate(updated_metadatas):
680                if target_key in metadata:
681                    continue
682
683                source_value = metadata.get(source_key)
684                if not isinstance(source_value, str):
685                    continue
686
687                inputs.append(source_value)
688                positions.append(idx)
689
690            if len(inputs) == 0:
691                continue
692
693            sparse_embeddings = self._sparse_embed(
694                input=inputs,
695                sparse_embedding_function=embedding_func,
696            )
697
698            if len(sparse_embeddings) != len(positions):
699                raise ValueError(
700                    "Sparse embedding function returned unexpected number of embeddings."
701                )
702
703            for position, embedding in zip(positions, sparse_embeddings):
704                updated_metadatas[position][target_key] = embedding
705
706        # Convert empty dicts back to None, validation requires non-empty dicts or None
707        result_metadatas: List[Optional[Metadata]] = [
708            metadata if metadata else None for metadata in updated_metadatas
709        ]
710
711        validate_metadatas(cast(List[Metadata], result_metadatas))
712        return cast(List[Metadata], result_metadatas)
713
714    def _embed_record_set(
715        self,
716        record_set: BaseRecordSet,
717        embeddable_fields: Optional[Set[str]] = None,
718        is_query: bool = False,
719    ) -> Embeddings:
720        if embeddable_fields is None:
721            embeddable_fields = get_default_embeddable_record_set_fields()
722
723        for field in embeddable_fields:
724            if record_set[field] is not None:  # type: ignore[literal-required]
725                # uris require special handling
726                if field == "uris":
727                    if self._data_loader is None:
728                        raise ValueError(
729                            "You must set a data loader on the collection if loading from URIs."
730                        )
731                    return self._embed(
732                        input=self._data_loader(uris=cast(URIs, record_set[field])),  # type: ignore[literal-required]
733                        is_query=is_query,
734                    )
735                else:
736                    return self._embed(
737                        input=record_set[field],  # type: ignore[literal-required]
738                        is_query=is_query,
739                    )
740        raise ValueError(
741            "Record does not contain any non-None fields that can be embedded."
742            f"Embeddable Fields: {embeddable_fields}"
743            f"Record Fields: {record_set}"
744        )
745
746    def _embed(self, input: Any, is_query: bool = False) -> Embeddings:
747        if self._embedding_function is not None and not isinstance(
748            self._embedding_function, DefaultEmbeddingFunction
749        ):
750            if is_query:
751                return self._embedding_function.embed_query(input=input)
752            else:
753                return self._embedding_function(input=input)
754
755        config_ef = self.configuration.get("embedding_function")
756        if config_ef is not None:
757            if is_query:
758                return config_ef.embed_query(input=input)
759            else:
760                return config_ef(input=input)
761        schema = self.schema
762        schema_embedding_function: Optional[EmbeddingFunction[Embeddable]] = None
763        if schema is not None:
764            override = schema.keys.get(EMBEDDING_KEY)
765            if (
766                override is not None
767                and override.float_list is not None
768                and override.float_list.vector_index is not None
769                and override.float_list.vector_index.config.embedding_function
770                is not None
771            ):
772                schema_embedding_function = cast(
773                    EmbeddingFunction[Embeddable],
774                    override.float_list.vector_index.config.embedding_function,
775                )
776            elif (
777                schema.defaults.float_list is not None
778                and schema.defaults.float_list.vector_index is not None
779                and schema.defaults.float_list.vector_index.config.embedding_function
780                is not None
781            ):
782                schema_embedding_function = cast(
783                    EmbeddingFunction[Embeddable],
784                    schema.defaults.float_list.vector_index.config.embedding_function,
785                )
786
787        if schema_embedding_function is not None:
788            if is_query and hasattr(schema_embedding_function, "embed_query"):
789                return schema_embedding_function.embed_query(input=input)
790            return schema_embedding_function(input=input)
791        if self._embedding_function is None:
792            raise ValueError(
793                "You must provide an embedding function to compute embeddings."
794                "https://docs.trychroma.com/guides/embeddings"
795            )
796        if is_query:
797            return self._embedding_function.embed_query(input=input)
798        else:
799            return self._embedding_function(input=input)
800
801    def _sparse_embed(
802        self,
803        input: Any,
804        sparse_embedding_function: SparseEmbeddingFunction[Any],
805        is_query: bool = False,
806    ) -> Any:
807        if is_query:
808            return sparse_embedding_function.embed_query(input=input)
809        return sparse_embedding_function(input=input)
810
811    def _embed_knn_string_queries(self, knn: Any) -> Any:
812        """Embed string queries in Knn objects using the appropriate embedding function.
813
814        Args:
815            knn: A Knn object that may have a string query
816
817        Returns:
818            A Knn object with the string query replaced by an embedding
819
820        Raises:
821            ValueError: If the query is a string but no embedding function is available
822        """
823        from chromadb.execution.expression.operator import Knn
824
825        if not isinstance(knn, Knn):
826            return knn
827
828        # If query is not a string, nothing to do
829        if not isinstance(knn.query, str):
830            return knn
831
832        query_text = knn.query
833        key = knn.key
834
835        # Handle main embedding field
836        if key == EMBEDDING_KEY:
837            # Use the collection's main embedding function
838            embedding = self._embed(input=[query_text], is_query=True)
839            if not embedding or len(embedding) != 1:
840                raise ValueError(
841                    "Embedding function returned unexpected number of embeddings"
842                )
843            # Return a new Knn with the embedded query
844            return Knn(
845                query=embedding[0],
846                key=knn.key,
847                limit=knn.limit,
848                default=knn.default,
849                return_rank=knn.return_rank,
850            )
851
852        # Handle metadata field with potential sparse embedding
853        schema = self.schema
854        if schema is None or key not in schema.keys:
855            raise ValueError(
856                f"Cannot embed string query for key '{key}': "
857                f"key not found in schema. Please provide an embedded vector or "
858                f"configure an embedding function for this key in the schema."
859            )
860
861        value_type = schema.keys[key]
862
863        # Check for sparse vector with embedding function
864        if value_type.sparse_vector is not None:
865            sparse_index = value_type.sparse_vector.sparse_vector_index
866            if sparse_index is not None and sparse_index.enabled:
867                sparse_config = sparse_index.config
868                if sparse_config.embedding_function is not None:
869                    embedding_func = sparse_config.embedding_function
870                    if not isinstance(embedding_func, SparseEmbeddingFunction):
871                        embedding_func = cast(
872                            SparseEmbeddingFunction[Any], embedding_func
873                        )
874                    validate_sparse_embedding_function(embedding_func)
875
876                    # Embed the query
877                    sparse_embedding = self._sparse_embed(
878                        input=[query_text],
879                        sparse_embedding_function=embedding_func,
880                        is_query=True,
881                    )
882
883                    if not sparse_embedding or len(sparse_embedding) != 1:
884                        raise ValueError(
885                            "Sparse embedding function returned unexpected number of embeddings"
886                        )
887
888                    # Return a new Knn with the sparse embedding
889                    return Knn(
890                        query=sparse_embedding[0],
891                        key=knn.key,
892                        limit=knn.limit,
893                        default=knn.default,
894                        return_rank=knn.return_rank,
895                    )
896
897        # Check for dense vector with embedding function (float_list)
898        if value_type.float_list is not None:
899            vector_index = value_type.float_list.vector_index
900            if vector_index is not None and vector_index.enabled:
901                dense_config = vector_index.config
902                if dense_config.embedding_function is not None:
903                    embedding_func = dense_config.embedding_function
904                    validate_embedding_function(embedding_func)
905
906                    # Embed the query using the schema's embedding function
907                    try:
908                        embeddings = embedding_func.embed_query(input=[query_text])
909                    except AttributeError:
910                        # Fallback if embed_query doesn't exist
911                        embeddings = embedding_func([query_text])
912
913                    if not embeddings or len(embeddings) != 1:
914                        raise ValueError(
915                            "Embedding function returned unexpected number of embeddings"
916                        )
917
918                    # Return a new Knn with the dense embedding
919                    return Knn(
920                        query=embeddings[0],
921                        key=knn.key,
922                        limit=knn.limit,
923                        default=knn.default,
924                        return_rank=knn.return_rank,
925                    )
926
927        raise ValueError(
928            f"Cannot embed string query for key '{key}': "
929            f"no embedding function configured for this key in the schema. "
930            f"Please provide an embedded vector or configure an embedding function."
931        )
932
933    def _embed_rank_string_queries(self, rank: Any) -> Any:
934        """Recursively embed string queries in Rank expressions.
935
936        Args:
937            rank: A Rank expression that may contain Knn objects with string queries
938
939        Returns:
940            A Rank expression with all string queries embedded
941        """
942        # Import here to avoid circular dependency
943        from chromadb.execution.expression.operator import (
944            Knn,
945            Abs,
946            Div,
947            Exp,
948            Log,
949            Max,
950            Min,
951            Mul,
952            Sub,
953            Sum,
954            Val,
955            Rrf,
956        )
957
958        if rank is None:
959            return None
960
961        # Base case: Knn - embed if it has a string query
962        if isinstance(rank, Knn):
963            return self._embed_knn_string_queries(rank)
964
965        # Base case: Val - no embedding needed
966        if isinstance(rank, Val):
967            return rank
968
969        # Recursive cases: walk through child ranks
970        if isinstance(rank, Abs):
971            return Abs(self._embed_rank_string_queries(rank.rank))
972
973        if isinstance(rank, Div):
974            return Div(
975                self._embed_rank_string_queries(rank.left),
976                self._embed_rank_string_queries(rank.right),
977            )
978
979        if isinstance(rank, Exp):
980            return Exp(self._embed_rank_string_queries(rank.rank))
981
982        if isinstance(rank, Log):
983            return Log(self._embed_rank_string_queries(rank.rank))
984
985        if isinstance(rank, Max):
986            return Max([self._embed_rank_string_queries(r) for r in rank.ranks])
987
988        if isinstance(rank, Min):
989            return Min([self._embed_rank_string_queries(r) for r in rank.ranks])
990
991        if isinstance(rank, Mul):
992            return Mul([self._embed_rank_string_queries(r) for r in rank.ranks])
993
994        if isinstance(rank, Sub):
995            return Sub(
996                self._embed_rank_string_queries(rank.left),
997                self._embed_rank_string_queries(rank.right),
998            )
999
1000        if isinstance(rank, Sum):
1001            return Sum([self._embed_rank_string_queries(r) for r in rank.ranks])
1002
1003        if isinstance(rank, Rrf):
1004            return Rrf(
1005                ranks=[self._embed_rank_string_queries(r) for r in rank.ranks],
1006                k=rank.k,
1007                weights=rank.weights,
1008                normalize=rank.normalize,
1009            )
1010
1011        # Unknown rank type - return as is
1012        return rank
1013
1014    def _embed_search_string_queries(self, search: Any) -> Any:
1015        """Embed string queries in a Search object.
1016
1017        Args:
1018            search: A Search object that may contain Knn objects with string queries
1019
1020        Returns:
1021            A Search object with all string queries embedded
1022        """
1023        # Import here to avoid circular dependency
1024        from chromadb.execution.expression.plan import Search
1025
1026        if not isinstance(search, Search):
1027            return search
1028
1029        # Embed the rank expression if it exists
1030        embedded_rank = self._embed_rank_string_queries(search._rank)
1031
1032        # Create a new Search with the embedded rank
1033        return Search(
1034            where=search._where,
1035            rank=embedded_rank,
1036            group_by=search._group_by,
1037            limit=search._limit,
1038            select=search._select,
1039        )
1040 
codekingpro/portable-devtools · Team Ai