Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
Collection.py664 linesDownload Raw Back to models
1from typing import TYPE_CHECKING, Optional, Union, List, cast, Dict, Any, Tuple
2
3from chromadb.api.models.CollectionCommon import CollectionCommon
4from chromadb.api.types import (
5    URI,
6    CollectionMetadata,
7    Embedding,
8    PyEmbedding,
9    Include,
10    IndexingStatus,
11    Metadata,
12    Document,
13    Image,
14    Where,
15    IDs,
16    GetResult,
17    QueryResult,
18    ID,
19    OneOrMany,
20    ReadLevel,
21    WhereDocument,
22    SearchResult,
23    DeleteResult,
24    maybe_cast_one_to_many,
25)
26from chromadb.api.collection_configuration import UpdateCollectionConfiguration
27from chromadb.execution.expression.plan import Search
28
29import logging
30
31from chromadb.api.functions import Function
32
33if TYPE_CHECKING:
34    from chromadb.api.models.AttachedFunction import AttachedFunction
35
36logger = logging.getLogger(__name__)
37
38if TYPE_CHECKING:
39    from chromadb.api import ServerAPI  # noqa: F401
40
41
42class Collection(CollectionCommon["ServerAPI"]):
43    def count(self, read_level: ReadLevel = ReadLevel.INDEX_AND_WAL) -> int:
44        """Return the number of records in the collection.
45
46        Args:
47            read_level: Controls whether to read from the write-ahead log (WAL):
48                - ReadLevel.INDEX_AND_WAL: Read from both the compacted index and WAL (default).
49                  All committed writes will be visible.
50                - ReadLevel.INDEX_ONLY: Read only from the compacted index, skipping the WAL.
51                  Faster, but recent writes that haven't been compacted may not be visible.
52                - ReadLevel.INDEX_AND_BOUNDED_WAL: Read from the index and up to a
53                  server-configured number of WAL entries for bounded query latency.
54        """
55        return self._client._count(
56            collection_id=self.id,
57            tenant=self.tenant,
58            database=self.database,
59            read_level=read_level,
60        )
61
62    def get_indexing_status(self) -> IndexingStatus:
63        """Get the indexing status of this collection.
64
65        Returns:
66            IndexingStatus: An object containing:
67                - num_indexed_ops: Number of user operations that have been indexed
68                - num_unindexed_ops: Number of user operations pending indexing
69                - total_ops: Total number of user operations in collection
70                - op_indexing_progress: Proportion of user operations that have been indexed as a float between 0 and 1
71        """
72        return self._client._get_indexing_status(
73            collection_id=self.id,
74            tenant=self.tenant,
75            database=self.database,
76        )
77
78    def add(
79        self,
80        ids: OneOrMany[ID],
81        embeddings: Optional[
82            Union[
83                OneOrMany[Embedding],
84                OneOrMany[PyEmbedding],
85            ]
86        ] = None,
87        metadatas: Optional[OneOrMany[Metadata]] = None,
88        documents: Optional[OneOrMany[Document]] = None,
89        images: Optional[OneOrMany[Image]] = None,
90        uris: Optional[OneOrMany[URI]] = None,
91    ) -> None:
92        """Add records to the collection.
93
94        Args:
95            ids: Record IDs to add.
96            embeddings: Embeddings to add. If None, embeddings are computed.
97            metadatas: Optional metadata for each record.
98            documents: Optional documents for each record.
99            images: Optional images for each record.
100            uris: Optional URIs for loading images.
101
102        Raises:
103            ValueError: If embeddings and documents are both missing.
104            ValueError: If embeddings and documents are both provided.
105            ValueError: If lengths of provided fields do not match.
106            ValueError: If an ID already exists.
107        """
108
109        add_request = self._validate_and_prepare_add_request(
110            ids=ids,
111            embeddings=embeddings,
112            metadatas=metadatas,
113            documents=documents,
114            images=images,
115            uris=uris,
116        )
117
118        self._client._add(
119            collection_id=self.id,
120            ids=add_request["ids"],
121            embeddings=add_request["embeddings"],
122            metadatas=add_request["metadatas"],
123            documents=add_request["documents"],
124            uris=add_request["uris"],
125            tenant=self.tenant,
126            database=self.database,
127        )
128
129    def get(
130        self,
131        ids: Optional[OneOrMany[ID]] = None,
132        where: Optional[Where] = None,
133        limit: Optional[int] = None,
134        offset: Optional[int] = None,
135        where_document: Optional[WhereDocument] = None,
136        include: Include = ["metadatas", "documents"],
137    ) -> GetResult:
138        """Retrieve records from the collection.
139
140        If no filters are provided, returns records up to ``limit`` starting at
141        ``offset``.
142
143        Args:
144            ids: If provided, only return records with these IDs.
145            where: A Where filter used to filter based on metadata values.
146            limit: Maximum number of results to return.
147            offset: Number of results to skip before returning.
148            where_document: A WhereDocument filter used to filter based on K.DOCUMENT.
149            include: Fields to include in results. Can contain "embeddings", "metadatas", "documents", "uris". Defaults to "metadatas" and "documents".
150
151        Returns:
152            GetResult: Retrieved records and requested fields as a GetResult object.
153        """
154        get_request = self._validate_and_prepare_get_request(
155            ids=ids,
156            where=where,
157            where_document=where_document,
158            include=include,
159        )
160
161        get_results = self._client._get(
162            collection_id=self.id,
163            ids=get_request["ids"],
164            where=get_request["where"],
165            where_document=get_request["where_document"],
166            include=get_request["include"],
167            limit=limit,
168            offset=offset,
169            tenant=self.tenant,
170            database=self.database,
171        )
172        return self._transform_get_response(
173            response=get_results, include=get_request["include"]
174        )
175
176    def peek(self, limit: int = 10) -> GetResult:
177        """Return the first ``limit`` records from the collection.
178
179        Args:
180            limit: Maximum number of records to return.
181
182        Returns:
183            GetResult: Retrieved records and requested fields.
184        """
185        return self._transform_peek_response(
186            self._client._peek(
187                collection_id=self.id,
188                n=limit,
189                tenant=self.tenant,
190                database=self.database,
191            )
192        )
193
194    def query(
195        self,
196        query_embeddings: Optional[
197            Union[
198                OneOrMany[Embedding],
199                OneOrMany[PyEmbedding],
200            ]
201        ] = None,
202        query_texts: Optional[OneOrMany[Document]] = None,
203        query_images: Optional[OneOrMany[Image]] = None,
204        query_uris: Optional[OneOrMany[URI]] = None,
205        ids: Optional[OneOrMany[ID]] = None,
206        n_results: int = 10,
207        where: Optional[Where] = None,
208        where_document: Optional[WhereDocument] = None,
209        include: Include = [
210            "metadatas",
211            "documents",
212            "distances",
213        ],
214    ) -> QueryResult:
215        """Query for the K nearest neighbor records in the collection.
216
217        This is a batch query API. Multiple queries can be performed at once
218        by providing multiple embeddings, texts, or images.
219
220        >>> query_1 = [0.1, 0.2, 0.3]
221        >>> query_2 = [0.4, 0.5, 0.6]
222        >>> results = collection.query(
223        >>>     query_embeddings=[query_1, query_2],
224        >>>     n_results=10,
225        >>> )
226
227        If query_texts, query_images, or query_uris are provided, the collection's
228        embedding function will be used to create embeddings before querying
229        the API.
230
231        The `ids`, `where`, `where_document`, and `include` parameters are applied
232        to all queries.
233
234        Args:
235            query_embeddings: Raw embeddings to query for.
236            query_texts: Documents to embed and query against.
237            query_images: Images to embed and query against.
238            query_uris: URIs to be loaded and embedded.
239            ids: Optional subset of IDs to search within.
240            n_results: Number of neighbors to return per query.
241            where: Metadata filter.
242            where_document: Document content filter.
243            include: Fields to include in results. Can contain "embeddings", "metadatas", "documents", "uris", "distances". Defaults to "metadatas", "documents", "distances".
244
245        Returns:
246            QueryResult: Nearest neighbor results.
247
248        Raises:
249            ValueError: If no query input is provided.
250            ValueError: If multiple query input types are provided.
251        """
252
253        query_request = self._validate_and_prepare_query_request(
254            query_embeddings=query_embeddings,
255            query_texts=query_texts,
256            query_images=query_images,
257            query_uris=query_uris,
258            ids=ids,
259            n_results=n_results,
260            where=where,
261            where_document=where_document,
262            include=include,
263        )
264
265        query_results = self._client._query(
266            collection_id=self.id,
267            ids=query_request["ids"],
268            query_embeddings=query_request["embeddings"],
269            n_results=query_request["n_results"],
270            where=query_request["where"],
271            where_document=query_request["where_document"],
272            include=query_request["include"],
273            tenant=self.tenant,
274            database=self.database,
275        )
276
277        return self._transform_query_response(
278            response=query_results, include=query_request["include"]
279        )
280
281    def modify(
282        self,
283        name: Optional[str] = None,
284        metadata: Optional[CollectionMetadata] = None,
285        configuration: Optional[UpdateCollectionConfiguration] = None,
286    ) -> None:
287        """Update collection name, metadata, or configuration.
288
289        Args:
290            name: New collection name.
291            metadata: New metadata for the collection.
292            configuration: New configuration for the collection.
293        """
294
295        self._validate_modify_request(metadata)
296
297        # Note there is a race condition here where the metadata can be updated
298        # but another thread sees the cached local metadata.
299        # TODO: fixme
300        self._client._modify(
301            id=self.id,
302            new_name=name,
303            new_metadata=metadata,
304            new_configuration=configuration,
305            tenant=self.tenant,
306            database=self.database,
307        )
308
309        self._update_model_after_modify_success(name, metadata, configuration)
310
311    def fork(
312        self,
313        new_name: str,
314    ) -> "Collection":
315        """Fork the current collection under a new name. The returning collection should contain identical data to the current collection.
316        This only works for Hosted Chroma for now.
317
318        Args:
319            new_name: The name of the new collection.
320
321        Returns:
322            Collection: A new collection with the specified name and containing identical data to the current collection.
323        """
324        model = self._client._fork(
325            collection_id=self.id,
326            new_name=new_name,
327            tenant=self.tenant,
328            database=self.database,
329        )
330        return Collection(
331            client=self._client,
332            model=model,
333            embedding_function=self._embedding_function,
334            data_loader=self._data_loader,
335        )
336
337    def fork_count(self) -> int:
338        """Get the number of forks that exist for this collection.
339        This only works for Hosted Chroma for now.
340
341        Returns:
342            int: The number of forks for this collection.
343        """
344        return self._client._fork_count(
345            collection_id=self.id,
346            tenant=self.tenant,
347            database=self.database,
348        )
349
350    def search(
351        self,
352        searches: OneOrMany[Search],
353        read_level: ReadLevel = ReadLevel.INDEX_AND_WAL,
354    ) -> SearchResult:
355        """Perform hybrid search on the collection.
356        This is an experimental API that only works for distributed and hosted Chroma for now.
357
358        Args:
359            searches: A single Search object or a list of Search objects, each containing:
360                - where: Where expression for filtering
361                - rank: Ranking expression for hybrid search (defaults to Val(0.0))
362                - limit: Limit configuration for pagination (defaults to no limit)
363                - select: Select configuration for keys to return (defaults to empty)
364            read_level: Controls whether to read from the write-ahead log (WAL):
365                - ReadLevel.INDEX_AND_WAL: Read from both the compacted index and WAL (default).
366                  All committed writes will be visible.
367                - ReadLevel.INDEX_ONLY: Read only from the compacted index, skipping the WAL.
368                  Faster, but recent writes that haven't been compacted may not be visible.
369                - ReadLevel.INDEX_AND_BOUNDED_WAL: Read from the index and up to a
370                  server-configured number of WAL entries for bounded query latency.
371
372        Returns:
373            SearchResult: Column-major format response with:
374                - ids: List of result IDs for each search payload
375                - documents: Optional documents for each payload
376                - embeddings: Optional embeddings for each payload
377                - metadatas: Optional metadata for each payload
378                - scores: Optional scores for each payload
379                - select: List of selected keys for each payload
380
381        Raises:
382            NotImplementedError: For local/segment API implementations
383
384        Examples:
385            # Using builder pattern with Key constants
386            from chromadb.execution.expression import (
387                Search, Key, K, Knn, Val
388            )
389
390            # Note: K is an alias for Key, so K.DOCUMENT == Key.DOCUMENT
391            search = (Search()
392                .where((K("category") == "science") & (K("score") > 0.5))
393                .rank(Knn(query=[0.1, 0.2, 0.3]) * 0.8 + Val(0.5) * 0.2)
394                .limit(10, offset=0)
395                .select(K.DOCUMENT, K.SCORE, "title"))
396
397            # Direct construction
398            from chromadb.execution.expression import (
399                Search, Eq, And, Gt, Knn, Limit, Select, Key
400            )
401
402            search = Search(
403                where=And([Eq("category", "science"), Gt("score", 0.5)]),
404                rank=Knn(query=[0.1, 0.2, 0.3]),
405                limit=Limit(offset=0, limit=10),
406                select=Select(keys={Key.DOCUMENT, Key.SCORE, "title"})
407            )
408
409            # Single search
410            result = collection.search(search)
411
412            # Multiple searches at once
413            searches = [
414                Search().where(K("type") == "article").rank(Knn(query=[0.1, 0.2])),
415                Search().where(K("type") == "paper").rank(Knn(query=[0.3, 0.4]))
416            ]
417            results = collection.search(searches)
418
419            # Skip WAL for faster queries (may miss recent uncommitted writes)
420            from chromadb.api.types import ReadLevel
421            result = collection.search(search, read_level=ReadLevel.INDEX_ONLY)
422        """
423        # Convert single search to list for consistent handling
424        searches_list = maybe_cast_one_to_many(searches)
425        if searches_list is None:
426            searches_list = []
427
428        # Embed any string queries in Knn objects
429        embedded_searches = [
430            self._embed_search_string_queries(search) for search in searches_list
431        ]
432
433        return self._client._search(
434            collection_id=self.id,
435            searches=cast(List[Search], embedded_searches),
436            tenant=self.tenant,
437            database=self.database,
438            read_level=read_level,
439        )
440
441    def update(
442        self,
443        ids: OneOrMany[ID],
444        embeddings: Optional[
445            Union[
446                OneOrMany[Embedding],
447                OneOrMany[PyEmbedding],
448            ]
449        ] = None,
450        metadatas: Optional[OneOrMany[Metadata]] = None,
451        documents: Optional[OneOrMany[Document]] = None,
452        images: Optional[OneOrMany[Image]] = None,
453        uris: Optional[OneOrMany[URI]] = None,
454    ) -> None:
455        """Update existing records by ID.
456
457        Records are provided in columnar format. If provided, the `embeddings`, `metadatas`, `documents`, and `uris` lists must be the same length.
458        Entries in each list correspond to the same record.
459
460        >>> ids = ["id1", "id2", "id3"]
461        >>> embeddings = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6], [0.7, 0.8, 0.9]]
462        >>> metadatas = [{"key": "value"}, {"key": "value"}, {"key": "value"}]
463        >>> documents = ["document1", "document2", "document3"]
464        >>> uris = ["uri1", "uri2", "uri3"]
465        >>> collection.update(ids, embeddings, metadatas, documents, uris)
466
467        If `embeddings` are not provided, the embeddings will be computed based on `documents` using the collection's embedding function.
468
469        Args:
470            ids: Record IDs to update.
471            embeddings: Updated embeddings. If None, embeddings are computed.
472            metadatas: Updated metadata.
473            documents: Updated documents.
474            images: Updated images.
475            uris: Updated URIs for loading images.
476        """
477        update_request = self._validate_and_prepare_update_request(
478            ids=ids,
479            embeddings=embeddings,
480            metadatas=metadatas,
481            documents=documents,
482            images=images,
483            uris=uris,
484        )
485
486        self._client._update(
487            collection_id=self.id,
488            ids=update_request["ids"],
489            embeddings=update_request["embeddings"],
490            metadatas=update_request["metadatas"],
491            documents=update_request["documents"],
492            uris=update_request["uris"],
493            tenant=self.tenant,
494            database=self.database,
495        )
496
497    def upsert(
498        self,
499        ids: OneOrMany[ID],
500        embeddings: Optional[
501            Union[
502                OneOrMany[Embedding],
503                OneOrMany[PyEmbedding],
504            ]
505        ] = None,
506        metadatas: Optional[OneOrMany[Metadata]] = None,
507        documents: Optional[OneOrMany[Document]] = None,
508        images: Optional[OneOrMany[Image]] = None,
509        uris: Optional[OneOrMany[URI]] = None,
510    ) -> None:
511        """Create or update records by ID.
512
513        Args:
514            ids: Record IDs to upsert.
515            embeddings: Embeddings to add or update. If None, embeddings are computed.
516            metadatas: Metadata to add or update.
517            documents: Documents to add or update.
518            images: Images to add or update.
519            uris: URIs for loading images.
520        """
521        upsert_request = self._validate_and_prepare_upsert_request(
522            ids=ids,
523            embeddings=embeddings,
524            metadatas=metadatas,
525            documents=documents,
526            images=images,
527            uris=uris,
528        )
529
530        self._client._upsert(
531            collection_id=self.id,
532            ids=upsert_request["ids"],
533            embeddings=upsert_request["embeddings"],
534            metadatas=upsert_request["metadatas"],
535            documents=upsert_request["documents"],
536            uris=upsert_request["uris"],
537            tenant=self.tenant,
538            database=self.database,
539        )
540
541    def delete(
542        self,
543        ids: Optional[IDs] = None,
544        where: Optional[Where] = None,
545        where_document: Optional[WhereDocument] = None,
546        limit: Optional[int] = None,
547    ) -> DeleteResult:
548        """Delete records by ID or filters.
549
550        All documents that match the `ids` or `where` and `where_document` filters will be deleted.
551
552        Args:
553            ids: Record IDs to delete.
554            where: Metadata filter.
555            where_document: Document content filter.
556            limit: Maximum number of records to delete. Can only be used with where or where_document filters.
557
558        Returns:
559            DeleteResult: A dict containing the number of records deleted.
560
561        Raises:
562            ValueError: If no IDs or filters are provided.
563            ValueError: If limit is specified without a where or where_document clause.
564        """
565        delete_request = self._validate_and_prepare_delete_request(
566            ids, where, where_document, limit=limit
567        )
568
569        return self._client._delete(
570            collection_id=self.id,
571            ids=delete_request["ids"],
572            where=delete_request["where"],
573            where_document=delete_request["where_document"],
574            limit=delete_request["limit"],
575            tenant=self.tenant,
576            database=self.database,
577        )
578
579    def attach_function(
580        self,
581        function: Function,
582        name: str,
583        output_collection: str,
584        params: Optional[Dict[str, Any]] = None,
585    ) -> Tuple["AttachedFunction", bool]:
586        """Attach a function to this collection.
587
588        Args:
589            function: A Function enum value (e.g., STATISTICS_FUNCTION, RECORD_COUNTER_FUNCTION)
590            name: Unique name for this attached function
591            output_collection: Name of the collection where function output will be stored
592            params: Optional dictionary with function-specific parameters
593
594        Returns:
595            Tuple of (AttachedFunction, created) where created is True if newly created,
596            False if already existed (idempotent request)
597
598        Example:
599            >>> from chromadb.api.functions import STATISTICS_FUNCTION
600            >>> attached_fn = collection.attach_function(
601            ...     function=STATISTICS_FUNCTION,
602            ...     name="mycoll_stats_fn",
603            ...     output_collection="mycoll_stats",
604            ... )
605            >>> if created:
606            ...     print("New function attached")
607            ... else:
608            ...     print("Function already existed")
609        """
610        function_id = function.value if isinstance(function, Function) else function
611        return self._client.attach_function(
612            function_id=function_id,
613            name=name,
614            input_collection_id=self.id,
615            output_collection=output_collection,
616            params=params,
617            tenant=self.tenant,
618            database=self.database,
619        )
620
621    def get_attached_function(self, name: str) -> "AttachedFunction":
622        """Get an attached function by name for this collection.
623
624        Args:
625            name: Name of the attached function
626
627        Returns:
628            AttachedFunction: The attached function object
629
630        Raises:
631            NotFoundError: If the attached function doesn't exist
632        """
633        return self._client.get_attached_function(
634            name=name,
635            input_collection_id=self.id,
636            tenant=self.tenant,
637            database=self.database,
638        )
639
640    def detach_function(
641        self,
642        name: str,
643        delete_output_collection: bool = False,
644    ) -> bool:
645        """Detach a function from this collection.
646
647        Args:
648            name: The name of the attached function
649            delete_output_collection: Whether to also delete the output collection. Defaults to False.
650
651        Returns:
652            bool: True if successful
653
654        Example:
655            >>> success = collection.detach_function("my_function", delete_output_collection=True)
656        """
657        return self._client.detach_function(
658            name=name,
659            input_collection_id=self.id,
660            delete_output=delete_output_collection,
661            tenant=self.tenant,
662            database=self.database,
663        )
664 
codekingpro/portable-devtools · Team Ai