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