Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
segment.py1128 linesDownload Raw Back to api
1from typing import TYPE_CHECKING
2
3from tenacity import retry, stop_after_attempt, retry_if_exception, wait_fixed
4from chromadb.api import ServerAPI
5
6if TYPE_CHECKING:
7    from chromadb.api.models.AttachedFunction import AttachedFunction
8from chromadb.api.collection_configuration import (
9    CreateCollectionConfiguration,
10    UpdateCollectionConfiguration,
11    create_collection_configuration_to_json,
12)
13from chromadb.auth import UserIdentity
14from chromadb.config import DEFAULT_DATABASE, DEFAULT_TENANT, Settings, System
15from chromadb.db.system import SysDB
16from chromadb.quota import QuotaEnforcer, Action
17from chromadb.rate_limit import RateLimitEnforcer
18from chromadb.segment import SegmentManager
19from chromadb.execution.executor.abstract import Executor
20from chromadb.execution.expression.operator import Scan, Filter, Limit, KNN, Projection
21from chromadb.execution.expression.plan import CountPlan, GetPlan, KNNPlan
22from chromadb.telemetry.opentelemetry import (
23    add_attributes_to_current_span,
24    OpenTelemetryClient,
25    OpenTelemetryGranularity,
26    trace_method,
27)
28from chromadb.telemetry.product import ProductTelemetryClient
29from chromadb.ingest import Producer
30from chromadb.types import Collection as CollectionModel
31from chromadb import __version__
32from chromadb.errors import (
33    InvalidDimensionException,
34    NotFoundError,
35    VersionMismatchError,
36)
37from chromadb.api.types import (
38    CollectionMetadata,
39    IDs,
40    Embeddings,
41    Metadatas,
42    Documents,
43    ReadLevel,
44    Schema,
45    URIs,
46    Where,
47    WhereDocument,
48    Include,
49    GetResult,
50    QueryResult,
51    SearchResult,
52    validate_metadata,
53    validate_update_metadata,
54    validate_where,
55    validate_where_document,
56    validate_batch,
57    IncludeMetadataDocuments,
58    IncludeMetadataDocumentsDistances,
59    DeleteResult,
60)
61from chromadb.telemetry.product.events import (
62    CollectionAddEvent,
63    CollectionDeleteEvent,
64    CollectionGetEvent,
65    CollectionUpdateEvent,
66    CollectionQueryEvent,
67    ClientCreateCollectionEvent,
68)
69
70import chromadb.types as t
71from typing import (
72    Optional,
73    Sequence,
74    Generator,
75    List,
76    Any,
77    Dict,
78    Callable,
79    TypeVar,
80    Tuple,
81)
82from overrides import override
83from uuid import UUID, uuid4
84from functools import wraps
85import time
86import logging
87import re
88from chromadb.execution.expression.plan import Search
89
90T = TypeVar("T", bound=Callable[..., Any])
91
92logger = logging.getLogger(__name__)
93
94
95# mimics s3 bucket requirements for naming
96def check_index_name(index_name: str) -> None:
97    msg = (
98        "Expected collection name that "
99        "(1) contains 3-63 characters, "
100        "(2) starts and ends with an alphanumeric character, "
101        "(3) otherwise contains only alphanumeric characters, underscores or hyphens (-), "
102        "(4) contains no two consecutive periods (..) and "
103        "(5) is not a valid IPv4 address, "
104        f"got {index_name}"
105    )
106    if len(index_name) < 3 or len(index_name) > 63:
107        raise ValueError(msg)
108    if not re.match("^[a-zA-Z0-9][a-zA-Z0-9._-]*[a-zA-Z0-9]$", index_name):
109        raise ValueError(msg)
110    if ".." in index_name:
111        raise ValueError(msg)
112    if re.match("^[0-9]{1,3}\\.[0-9]{1,3}\\.[0-9]{1,3}\\.[0-9]{1,3}$", index_name):
113        raise ValueError(msg)
114
115
116def rate_limit(func: T) -> T:
117    @wraps(func)
118    def wrapper(*args: Any, **kwargs: Any) -> Any:
119        self = args[0]
120        return self._rate_limit_enforcer.rate_limit(func)(*args, **kwargs)
121
122    return wrapper  # type: ignore
123
124
125class SegmentAPI(ServerAPI):
126    """API implementation utilizing the new segment-based internal architecture"""
127
128    _settings: Settings
129    _sysdb: SysDB
130    _manager: SegmentManager
131    _executor: Executor
132    _producer: Producer
133    _product_telemetry_client: ProductTelemetryClient
134    _opentelemetry_client: OpenTelemetryClient
135    _tenant_id: str
136    _topic_ns: str
137    _rate_limit_enforcer: RateLimitEnforcer
138
139    def __init__(self, system: System):
140        super().__init__(system)
141        self._settings = system.settings
142        self._sysdb = self.require(SysDB)
143        self._manager = self.require(SegmentManager)
144        self._executor = self.require(Executor)
145        self._quota_enforcer = self.require(QuotaEnforcer)
146        self._product_telemetry_client = self.require(ProductTelemetryClient)
147        self._opentelemetry_client = self.require(OpenTelemetryClient)
148        self._producer = self.require(Producer)
149        self._rate_limit_enforcer = self._system.require(RateLimitEnforcer)
150
151    @override
152    def heartbeat(self) -> int:
153        return int(time.time_ns())
154
155    @trace_method("SegmentAPI.create_database", OpenTelemetryGranularity.OPERATION)
156    @override
157    def create_database(self, name: str, tenant: str = DEFAULT_TENANT) -> None:
158        if len(name) < 3:
159            raise ValueError("Database name must be at least 3 characters long")
160
161        self._quota_enforcer.enforce(
162            action=Action.CREATE_DATABASE,
163            tenant=tenant,
164            name=name,
165        )
166
167        self._sysdb.create_database(
168            id=uuid4(),
169            name=name,
170            tenant=tenant,
171        )
172
173    @trace_method("SegmentAPI.get_database", OpenTelemetryGranularity.OPERATION)
174    @override
175    def get_database(self, name: str, tenant: str = DEFAULT_TENANT) -> t.Database:
176        return self._sysdb.get_database(name=name, tenant=tenant)
177
178    @trace_method("SegmentAPI.delete_database", OpenTelemetryGranularity.OPERATION)
179    @override
180    def delete_database(self, name: str, tenant: str = DEFAULT_TENANT) -> None:
181        self._sysdb.delete_database(name=name, tenant=tenant)
182
183    @trace_method("SegmentAPI.list_databases", OpenTelemetryGranularity.OPERATION)
184    @override
185    def list_databases(
186        self,
187        limit: Optional[int] = None,
188        offset: Optional[int] = None,
189        tenant: str = DEFAULT_TENANT,
190    ) -> Sequence[t.Database]:
191        return self._sysdb.list_databases(limit=limit, offset=offset, tenant=tenant)
192
193    @trace_method("SegmentAPI.create_tenant", OpenTelemetryGranularity.OPERATION)
194    @override
195    def create_tenant(self, name: str) -> None:
196        if len(name) < 3:
197            raise ValueError("Tenant name must be at least 3 characters long")
198
199        self._sysdb.create_tenant(
200            name=name,
201        )
202
203    @override
204    def get_user_identity(self) -> UserIdentity:
205        return UserIdentity(
206            user_id="",
207            tenant=DEFAULT_TENANT,
208            databases=[DEFAULT_DATABASE],
209        )
210
211    @trace_method("SegmentAPI.get_tenant", OpenTelemetryGranularity.OPERATION)
212    @override
213    def get_tenant(self, name: str) -> t.Tenant:
214        return self._sysdb.get_tenant(name=name)
215
216    # TODO: Actually fix CollectionMetadata type to remove type: ignore flags. This is
217    # necessary because changing the value type from `Any` to`` `Union[str, int, float]`
218    # causes the system to somehow convert all values to strings.
219    @trace_method("SegmentAPI.create_collection", OpenTelemetryGranularity.OPERATION)
220    @override
221    @rate_limit
222    def create_collection(
223        self,
224        name: str,
225        schema: Optional[Schema] = None,
226        configuration: Optional[CreateCollectionConfiguration] = None,
227        metadata: Optional[CollectionMetadata] = None,
228        get_or_create: bool = False,
229        tenant: str = DEFAULT_TENANT,
230        database: str = DEFAULT_DATABASE,
231    ) -> CollectionModel:
232        if metadata is not None:
233            validate_metadata(metadata)
234
235        # TODO: remove backwards compatibility in naming requirements
236        check_index_name(name)
237
238        self._quota_enforcer.enforce(
239            action=Action.CREATE_COLLECTION,
240            tenant=tenant,
241            name=name,
242            metadata=metadata,
243        )
244
245        id = uuid4()
246
247        model = CollectionModel(
248            id=id,
249            name=name,
250            metadata=metadata,
251            serialized_schema=None,
252            configuration_json=create_collection_configuration_to_json(
253                configuration or CreateCollectionConfiguration(), metadata
254            ),
255            tenant=tenant,
256            database=database,
257            dimension=None,
258        )
259
260        # TODO: Let sysdb create the collection directly from the model
261        coll, created = self._sysdb.create_collection(
262            id=model.id,
263            name=model.name,
264            schema=schema,
265            configuration=configuration or CreateCollectionConfiguration(),
266            segments=[],  # Passing empty till backend changes are deployed.
267            metadata=model.metadata,
268            dimension=None,  # This is lazily populated on the first add
269            get_or_create=get_or_create,
270            tenant=tenant,
271            database=database,
272        )
273
274        if created:
275            segments = self._manager.prepare_segments_for_new_collection(coll)
276            for segment in segments:
277                self._sysdb.create_segment(segment)
278        else:
279            logger.debug(
280                f"Collection {name} already exists, returning existing collection."
281            )
282
283        # TODO: This event doesn't capture the get_or_create case appropriately
284        # TODO: Re-enable embedding function tracking in create_collection
285        self._product_telemetry_client.capture(
286            ClientCreateCollectionEvent(
287                collection_uuid=str(id),
288                # embedding_function=embedding_function.__class__.__name__,
289            )
290        )
291        add_attributes_to_current_span({"collection_uuid": str(id)})
292
293        return coll
294
295    @trace_method(
296        "SegmentAPI.get_or_create_collection", OpenTelemetryGranularity.OPERATION
297    )
298    @override
299    @rate_limit
300    def get_or_create_collection(
301        self,
302        name: str,
303        schema: Optional[Schema] = None,
304        configuration: Optional[CreateCollectionConfiguration] = None,
305        metadata: Optional[CollectionMetadata] = None,
306        tenant: str = DEFAULT_TENANT,
307        database: str = DEFAULT_DATABASE,
308    ) -> CollectionModel:
309        return self.create_collection(
310            name=name,
311            schema=schema,
312            metadata=metadata,
313            configuration=configuration,
314            get_or_create=True,
315            tenant=tenant,
316            database=database,
317        )
318
319    # TODO: Actually fix CollectionMetadata type to remove type: ignore flags. This is
320    # necessary because changing the value type from `Any` to`` `Union[str, int, float]`
321    # causes the system to somehow convert all values to strings
322    @trace_method("SegmentAPI.get_collection", OpenTelemetryGranularity.OPERATION)
323    @override
324    @rate_limit
325    def get_collection(
326        self,
327        name: Optional[str] = None,
328        tenant: str = DEFAULT_TENANT,
329        database: str = DEFAULT_DATABASE,
330    ) -> CollectionModel:
331        existing = self._sysdb.get_collections(
332            name=name, tenant=tenant, database=database
333        )
334
335        if existing:
336            return existing[0]
337        else:
338            raise NotFoundError(f"Collection {name} does not exist.")
339
340    @trace_method("SegmentAPI.get_collection_by_id", OpenTelemetryGranularity.OPERATION)
341    @override
342    @rate_limit
343    def get_collection_by_id(
344        self,
345        collection_id: UUID,
346        tenant: str = DEFAULT_TENANT,
347        database: str = DEFAULT_DATABASE,
348    ) -> CollectionModel:
349        existing = self._sysdb.get_collections(
350            id=collection_id, tenant=tenant, database=database
351        )
352
353        if existing:
354            return existing[0]
355        else:
356            raise NotFoundError(f"Collection {collection_id} does not exist.")
357
358    @trace_method("SegmentAPI.list_collection", OpenTelemetryGranularity.OPERATION)
359    @override
360    @rate_limit
361    def list_collections(
362        self,
363        limit: Optional[int] = None,
364        offset: Optional[int] = None,
365        tenant: str = DEFAULT_TENANT,
366        database: str = DEFAULT_DATABASE,
367    ) -> Sequence[CollectionModel]:
368        self._quota_enforcer.enforce(
369            action=Action.LIST_COLLECTIONS,
370            tenant=tenant,
371            limit=limit,
372        )
373
374        return self._sysdb.get_collections(
375            limit=limit, offset=offset, tenant=tenant, database=database
376        )
377
378    @trace_method("SegmentAPI.count_collections", OpenTelemetryGranularity.OPERATION)
379    @override
380    @rate_limit
381    def count_collections(
382        self,
383        tenant: str = DEFAULT_TENANT,
384        database: str = DEFAULT_DATABASE,
385    ) -> int:
386        return self._sysdb.count_collections(tenant=tenant, database=database)
387
388    @trace_method("SegmentAPI._modify", OpenTelemetryGranularity.OPERATION)
389    @override
390    @rate_limit
391    def _modify(
392        self,
393        id: UUID,
394        new_name: Optional[str] = None,
395        new_metadata: Optional[CollectionMetadata] = None,
396        new_configuration: Optional[UpdateCollectionConfiguration] = None,
397        tenant: str = DEFAULT_TENANT,
398        database: str = DEFAULT_DATABASE,
399    ) -> None:
400        if new_name:
401            # backwards compatibility in naming requirements (for now)
402            check_index_name(new_name)
403
404        if new_metadata:
405            validate_update_metadata(new_metadata)
406
407        # Ensure the collection exists
408        _ = self._get_collection(id)
409
410        self._quota_enforcer.enforce(
411            action=Action.UPDATE_COLLECTION,
412            tenant=tenant,
413            name=new_name,
414            metadata=new_metadata,
415        )
416
417        # TODO eventually we'll want to use OptionalArgument and Unspecified in the
418        # signature of `_modify` but not changing the API right now.
419        if new_name and new_metadata and new_configuration:
420            self._sysdb.update_collection(
421                id,
422                name=new_name,
423                metadata=new_metadata,
424                configuration=new_configuration,
425            )
426        elif new_name and new_metadata:
427            self._sysdb.update_collection(id, name=new_name, metadata=new_metadata)
428        elif new_name and new_configuration:
429            self._sysdb.update_collection(
430                id, name=new_name, configuration=new_configuration
431            )
432        elif new_metadata and new_configuration:
433            self._sysdb.update_collection(
434                id, metadata=new_metadata, configuration=new_configuration
435            )
436        elif new_name:
437            self._sysdb.update_collection(id, name=new_name)
438        elif new_metadata:
439            self._sysdb.update_collection(id, metadata=new_metadata)
440        elif new_configuration:
441            self._sysdb.update_collection(id, configuration=new_configuration)
442
443    @override
444    def _fork(
445        self,
446        collection_id: UUID,
447        new_name: str,
448        tenant: str = DEFAULT_TENANT,
449        database: str = DEFAULT_DATABASE,
450    ) -> CollectionModel:
451        raise NotImplementedError(
452            "Collection forking is not implemented for SegmentAPI"
453        )
454
455    @override
456    def _fork_count(
457        self,
458        collection_id: UUID,
459        tenant: str = DEFAULT_TENANT,
460        database: str = DEFAULT_DATABASE,
461    ) -> int:
462        raise NotImplementedError(
463            "Fork count is not implemented for SegmentAPI"
464        )
465
466    @override
467    def _get_indexing_status(
468        self,
469        collection_id: UUID,
470        tenant: str = DEFAULT_TENANT,
471        database: str = DEFAULT_DATABASE,
472    ) -> "IndexingStatus":
473        raise NotImplementedError("Indexing status is not implemented for SegmentAPI")
474
475    @override
476    def _search(
477        self,
478        collection_id: UUID,
479        searches: List[Search],
480        tenant: str = DEFAULT_TENANT,
481        database: str = DEFAULT_DATABASE,
482        read_level: ReadLevel = ReadLevel.INDEX_AND_WAL,
483    ) -> SearchResult:
484        raise NotImplementedError("Search is not implemented for SegmentAPI")
485
486    @trace_method("SegmentAPI.delete_collection", OpenTelemetryGranularity.OPERATION)
487    @override
488    @rate_limit
489    def delete_collection(
490        self,
491        name: str,
492        tenant: str = DEFAULT_TENANT,
493        database: str = DEFAULT_DATABASE,
494    ) -> None:
495        existing = self._sysdb.get_collections(
496            name=name, tenant=tenant, database=database
497        )
498
499        if existing:
500            self._manager.delete_segments(existing[0].id)
501            self._sysdb.delete_collection(
502                existing[0].id, tenant=tenant, database=database
503            )
504        else:
505            raise ValueError(f"Collection {name} does not exist.")
506
507    @trace_method("SegmentAPI._add", OpenTelemetryGranularity.OPERATION)
508    @override
509    @rate_limit
510    def _add(
511        self,
512        ids: IDs,
513        collection_id: UUID,
514        embeddings: Embeddings,
515        metadatas: Optional[Metadatas] = None,
516        documents: Optional[Documents] = None,
517        uris: Optional[URIs] = None,
518        tenant: str = DEFAULT_TENANT,
519        database: str = DEFAULT_DATABASE,
520    ) -> bool:
521        coll = self._get_collection(collection_id)
522        self._manager.hint_use_collection(collection_id, t.Operation.ADD)
523        validate_batch(
524            (ids, embeddings, metadatas, documents, uris),
525            {"max_batch_size": self.get_max_batch_size()},
526        )
527        records_to_submit = list(
528            _records(
529                t.Operation.ADD,
530                ids=ids,
531                embeddings=embeddings,
532                metadatas=metadatas,
533                documents=documents,
534                uris=uris,
535            )
536        )
537        self._validate_embedding_record_set(coll, records_to_submit)
538
539        self._quota_enforcer.enforce(
540            action=Action.ADD,
541            tenant=tenant,
542            ids=ids,
543            embeddings=embeddings,
544            metadatas=metadatas,
545            documents=documents,
546            uris=uris,
547            collection_id=collection_id,
548        )
549
550        self._producer.submit_embeddings(collection_id, records_to_submit)
551
552        self._product_telemetry_client.capture(
553            CollectionAddEvent(
554                collection_uuid=str(collection_id),
555                add_amount=len(ids),
556                with_metadata=len(ids) if metadatas is not None else 0,
557                with_documents=len(ids) if documents is not None else 0,
558                with_uris=len(ids) if uris is not None else 0,
559            )
560        )
561        return True
562
563    @trace_method("SegmentAPI._update", OpenTelemetryGranularity.OPERATION)
564    @override
565    @rate_limit
566    def _update(
567        self,
568        collection_id: UUID,
569        ids: IDs,
570        embeddings: Optional[Embeddings] = None,
571        metadatas: Optional[Metadatas] = None,
572        documents: Optional[Documents] = None,
573        uris: Optional[URIs] = None,
574        tenant: str = DEFAULT_TENANT,
575        database: str = DEFAULT_DATABASE,
576    ) -> bool:
577        coll = self._get_collection(collection_id)
578        self._manager.hint_use_collection(collection_id, t.Operation.UPDATE)
579        validate_batch(
580            (ids, embeddings, metadatas, documents, uris),
581            {"max_batch_size": self.get_max_batch_size()},
582        )
583        records_to_submit = list(
584            _records(
585                t.Operation.UPDATE,
586                ids=ids,
587                embeddings=embeddings,
588                metadatas=metadatas,
589                documents=documents,
590                uris=uris,
591            )
592        )
593        self._validate_embedding_record_set(coll, records_to_submit)
594
595        self._quota_enforcer.enforce(
596            action=Action.UPDATE,
597            tenant=tenant,
598            ids=ids,
599            embeddings=embeddings,
600            metadatas=metadatas,
601            documents=documents,
602            uris=uris,
603        )
604
605        self._producer.submit_embeddings(collection_id, records_to_submit)
606
607        self._product_telemetry_client.capture(
608            CollectionUpdateEvent(
609                collection_uuid=str(collection_id),
610                update_amount=len(ids),
611                with_embeddings=len(embeddings) if embeddings else 0,
612                with_metadata=len(metadatas) if metadatas else 0,
613                with_documents=len(documents) if documents else 0,
614                with_uris=len(uris) if uris else 0,
615            )
616        )
617
618        return True
619
620    @trace_method("SegmentAPI._upsert", OpenTelemetryGranularity.OPERATION)
621    @override
622    @rate_limit
623    def _upsert(
624        self,
625        collection_id: UUID,
626        ids: IDs,
627        embeddings: Embeddings,
628        metadatas: Optional[Metadatas] = None,
629        documents: Optional[Documents] = None,
630        uris: Optional[URIs] = None,
631        tenant: str = DEFAULT_TENANT,
632        database: str = DEFAULT_DATABASE,
633    ) -> bool:
634        coll = self._get_collection(collection_id)
635        self._manager.hint_use_collection(collection_id, t.Operation.UPSERT)
636        validate_batch(
637            (ids, embeddings, metadatas, documents, uris),
638            {"max_batch_size": self.get_max_batch_size()},
639        )
640        records_to_submit = list(
641            _records(
642                t.Operation.UPSERT,
643                ids=ids,
644                embeddings=embeddings,
645                metadatas=metadatas,
646                documents=documents,
647                uris=uris,
648            )
649        )
650        self._validate_embedding_record_set(coll, records_to_submit)
651
652        self._quota_enforcer.enforce(
653            action=Action.UPSERT,
654            tenant=tenant,
655            ids=ids,
656            embeddings=embeddings,
657            metadatas=metadatas,
658            documents=documents,
659            uris=uris,
660            collection_id=collection_id,
661        )
662
663        self._producer.submit_embeddings(collection_id, records_to_submit)
664
665        return True
666
667    @trace_method("SegmentAPI._get", OpenTelemetryGranularity.OPERATION)
668    @retry(  # type: ignore[misc]
669        retry=retry_if_exception(lambda e: isinstance(e, VersionMismatchError)),
670        wait=wait_fixed(2),
671        stop=stop_after_attempt(5),
672        reraise=True,
673    )
674    @override
675    @rate_limit
676    def _get(
677        self,
678        collection_id: UUID,
679        ids: Optional[IDs] = None,
680        where: Optional[Where] = None,
681        limit: Optional[int] = None,
682        offset: Optional[int] = None,
683        where_document: Optional[WhereDocument] = None,
684        include: Include = IncludeMetadataDocuments,
685        tenant: str = DEFAULT_TENANT,
686        database: str = DEFAULT_DATABASE,
687    ) -> GetResult:
688        add_attributes_to_current_span(
689            {
690                "collection_id": str(collection_id),
691                "ids_count": len(ids) if ids else 0,
692            }
693        )
694
695        scan = self._scan(collection_id)
696
697        # TODO: Replace with unified validation
698        if where is not None:
699            validate_where(where)
700
701        if where_document is not None:
702            validate_where_document(where_document)
703
704        self._quota_enforcer.enforce(
705            action=Action.GET,
706            tenant=tenant,
707            ids=ids,
708            where=where,
709            where_document=where_document,
710            limit=limit,
711        )
712
713        ids_amount = len(ids) if ids else 0
714        self._product_telemetry_client.capture(
715            CollectionGetEvent(
716                collection_uuid=str(collection_id),
717                ids_count=ids_amount,
718                limit=limit if limit else 0,
719                include_metadata=ids_amount if "metadatas" in include else 0,
720                include_documents=ids_amount if "documents" in include else 0,
721                include_uris=ids_amount if "uris" in include else 0,
722            )
723        )
724
725        return self._executor.get(
726            GetPlan(
727                scan,
728                Filter(ids, where, where_document),
729                Limit(offset or 0, limit),
730                Projection(
731                    "documents" in include,
732                    "embeddings" in include,
733                    "metadatas" in include,
734                    False,
735                    "uris" in include,
736                ),
737            )
738        )
739
740    @trace_method("SegmentAPI._delete", OpenTelemetryGranularity.OPERATION)
741    @override
742    @rate_limit
743    def _delete(
744        self,
745        collection_id: UUID,
746        ids: Optional[IDs] = None,
747        where: Optional[Where] = None,
748        where_document: Optional[WhereDocument] = None,
749        limit: Optional[int] = None,
750        tenant: str = DEFAULT_TENANT,
751        database: str = DEFAULT_DATABASE,
752    ) -> DeleteResult:
753        add_attributes_to_current_span(
754            {
755                "collection_id": str(collection_id),
756                "ids_count": len(ids) if ids else 0,
757            }
758        )
759
760        # TODO: Replace with unified validation
761        if where is not None:
762            validate_where(where)
763
764        if where_document is not None:
765            validate_where_document(where_document)
766
767        # You must have at least one of non-empty ids, where, or where_document.
768        if (
769            (ids is None or (ids is not None and len(ids) == 0))
770            and (where is None or (where is not None and len(where) == 0))
771            and (
772                where_document is None
773                or (where_document is not None and len(where_document) == 0)
774            )
775        ):
776            raise ValueError(
777                """
778                You must provide either ids, where, or where_document to delete. If
779                you want to delete all data in a collection you can delete the
780                collection itself using the delete_collection method. Or alternatively,
781                you can get() all the relevant ids and then delete them.
782                """
783            )
784
785        scan = self._scan(collection_id)
786
787        self._quota_enforcer.enforce(
788            action=Action.DELETE,
789            tenant=tenant,
790            ids=ids,
791            where=where,
792            where_document=where_document,
793        )
794
795        self._manager.hint_use_collection(collection_id, t.Operation.DELETE)
796
797        if (where or where_document) or not ids:
798            ids_to_delete = self._executor.get(
799                GetPlan(scan, Filter(ids, where, where_document))
800            )["ids"]
801        else:
802            ids_to_delete = ids
803
804        # Apply limit if specified (validated upstream, but enforce defensively)
805        if limit is not None:
806            if not isinstance(limit, int) or isinstance(limit, bool) or limit < 0:
807                raise ValueError("limit must be a non-negative integer")
808            if where is None and where_document is None:
809                raise ValueError(
810                    "limit can only be specified when a where or where_document clause is provided"
811                )
812            ids_to_delete = ids_to_delete[:limit]
813
814        if len(ids_to_delete) == 0:
815            return DeleteResult(deleted=0)
816
817        records_to_submit = list(
818            _records(operation=t.Operation.DELETE, ids=ids_to_delete)
819        )
820        self._validate_embedding_record_set(scan.collection, records_to_submit)
821        self._producer.submit_embeddings(collection_id, records_to_submit)
822
823        deleted_count = len(ids_to_delete)
824
825        self._product_telemetry_client.capture(
826            CollectionDeleteEvent(
827                collection_uuid=str(collection_id), delete_amount=deleted_count
828            )
829        )
830
831        return DeleteResult(deleted=deleted_count)
832
833    @trace_method("SegmentAPI._count", OpenTelemetryGranularity.OPERATION)
834    @retry(  # type: ignore[misc]
835        retry=retry_if_exception(lambda e: isinstance(e, VersionMismatchError)),
836        wait=wait_fixed(2),
837        stop=stop_after_attempt(5),
838        reraise=True,
839    )
840    @override
841    @rate_limit
842    def _count(
843        self,
844        collection_id: UUID,
845        tenant: str = DEFAULT_TENANT,
846        database: str = DEFAULT_DATABASE,
847        read_level: ReadLevel = ReadLevel.INDEX_AND_WAL,
848    ) -> int:
849        add_attributes_to_current_span({"collection_id": str(collection_id)})
850        return self._executor.count(CountPlan(self._scan(collection_id)))
851
852    @trace_method("SegmentAPI._query", OpenTelemetryGranularity.OPERATION)
853    # We retry on version mismatch errors because the version of the collection
854    # may have changed between the time we got the version and the time we
855    # actually query the collection on the FE. We are fine with fixed
856    # wait time because the version mismatch error is not a error due to
857    # network issues or other transient issues. It is a result of the
858    # collection being updated between the time we got the version and
859    # the time we actually query the collection on the FE.
860    @retry(  # type: ignore[misc]
861        retry=retry_if_exception(lambda e: isinstance(e, VersionMismatchError)),
862        wait=wait_fixed(2),
863        stop=stop_after_attempt(5),
864        reraise=True,
865    )
866    @override
867    @rate_limit
868    def _query(
869        self,
870        collection_id: UUID,
871        query_embeddings: Embeddings,
872        ids: Optional[IDs] = None,
873        n_results: int = 10,
874        where: Optional[Where] = None,
875        where_document: Optional[WhereDocument] = None,
876        include: Include = IncludeMetadataDocumentsDistances,
877        tenant: str = DEFAULT_TENANT,
878        database: str = DEFAULT_DATABASE,
879    ) -> QueryResult:
880        add_attributes_to_current_span(
881            {
882                "collection_id": str(collection_id),
883                "n_results": n_results,
884                "where": str(where),
885            }
886        )
887
888        query_amount = len(query_embeddings)
889        ids_amount = len(ids) if ids else 0
890        self._product_telemetry_client.capture(
891            CollectionQueryEvent(
892                collection_uuid=str(collection_id),
893                query_amount=query_amount,
894                filtered_ids_amount=ids_amount,
895                n_results=n_results,
896                with_metadata_filter=query_amount if where is not None else 0,
897                with_document_filter=query_amount if where_document is not None else 0,
898                include_metadatas=query_amount if "metadatas" in include else 0,
899                include_documents=query_amount if "documents" in include else 0,
900                include_uris=query_amount if "uris" in include else 0,
901                include_distances=query_amount if "distances" in include else 0,
902            )
903        )
904
905        # TODO: Replace with unified validation
906        if where is not None:
907            validate_where(where)
908        if where_document is not None:
909            validate_where_document(where_document)
910
911        scan = self._scan(collection_id)
912        for embedding in query_embeddings:
913            self._validate_dimension(scan.collection, len(embedding), update=False)
914
915        self._quota_enforcer.enforce(
916            action=Action.QUERY,
917            tenant=tenant,
918            where=where,
919            where_document=where_document,
920            query_embeddings=query_embeddings,
921            n_results=n_results,
922        )
923
924        return self._executor.knn(
925            KNNPlan(
926                scan,
927                KNN(query_embeddings, n_results),
928                Filter(None, where, where_document),
929                Projection(
930                    "documents" in include,
931                    "embeddings" in include,
932                    "metadatas" in include,
933                    "distances" in include,
934                    "uris" in include,
935                ),
936            )
937        )
938
939    @trace_method("SegmentAPI._peek", OpenTelemetryGranularity.OPERATION)
940    @override
941    @rate_limit
942    def _peek(
943        self,
944        collection_id: UUID,
945        n: int = 10,
946        tenant: str = DEFAULT_TENANT,
947        database: str = DEFAULT_DATABASE,
948    ) -> GetResult:
949        add_attributes_to_current_span({"collection_id": str(collection_id)})
950        return self._get(collection_id, limit=n)  # type: ignore
951
952    @override
953    def get_version(self) -> str:
954        return __version__
955
956    @override
957    def reset_state(self) -> None:
958        pass
959
960    @override
961    def reset(self) -> bool:
962        self._system.reset_state()
963        return True
964
965    @override
966    def get_settings(self) -> Settings:
967        return self._settings
968
969    @override
970    def get_max_batch_size(self) -> int:
971        return self._producer.max_batch_size
972
973    @override
974    def attach_function(
975        self,
976        function_id: str,
977        name: str,
978        input_collection_id: UUID,
979        output_collection: str,
980        params: Optional[Dict[str, Any]] = None,
981        tenant: str = DEFAULT_TENANT,
982        database: str = DEFAULT_DATABASE,
983    ) -> Tuple["AttachedFunction", bool]:
984        """Attached functions are not supported in the Segment API (local embedded mode)."""
985        raise NotImplementedError(
986            "Attached functions are only supported when connecting to a Chroma server via HttpClient. "
987            "The Segment API (embedded mode) does not support attached function operations."
988        )
989
990    @override
991    def get_attached_function(
992        self,
993        name: str,
994        input_collection_id: UUID,
995        tenant: str = DEFAULT_TENANT,
996        database: str = DEFAULT_DATABASE,
997    ) -> "AttachedFunction":
998        """Attached functions are not supported in the Segment API (local embedded mode)."""
999        raise NotImplementedError(
1000            "Attached functions are only supported when connecting to a Chroma server via HttpClient. "
1001            "The Segment API (embedded mode) does not support attached function operations."
1002        )
1003
1004    @override
1005    def detach_function(
1006        self,
1007        name: str,
1008        input_collection_id: UUID,
1009        delete_output: bool = False,
1010        tenant: str = DEFAULT_TENANT,
1011        database: str = DEFAULT_DATABASE,
1012    ) -> bool:
1013        """Attached functions are not supported in the Segment API (local embedded mode)."""
1014        raise NotImplementedError(
1015            "Attached functions are only supported when connecting to a Chroma server via HttpClient. "
1016            "The Segment API (embedded mode) does not support attached function operations."
1017        )
1018
1019    # TODO: This could potentially cause race conditions in a distributed version of the
1020    # system, since the cache is only local.
1021    # TODO: promote collection -> topic to a base class method so that it can be
1022    # used for channel assignment in the distributed version of the system.
1023    @trace_method(
1024        "SegmentAPI._validate_embedding_record_set", OpenTelemetryGranularity.ALL
1025    )
1026    def _validate_embedding_record_set(
1027        self, collection: t.Collection, records: List[t.OperationRecord]
1028    ) -> None:
1029        """Validate the dimension of an embedding record before submitting it to the system."""
1030        add_attributes_to_current_span({"collection_id": str(collection["id"])})
1031        for record in records:
1032            if record["embedding"] is not None:
1033                self._validate_dimension(
1034                    collection, len(record["embedding"]), update=True
1035                )
1036
1037    # This method is intentionally left untraced because otherwise it can emit thousands of spans for requests containing many embeddings.
1038    def _validate_dimension(
1039        self, collection: t.Collection, dim: int, update: bool
1040    ) -> None:
1041        """Validate that a collection supports records of the given dimension. If update
1042        is true, update the collection if the collection doesn't already have a
1043        dimension."""
1044        if collection["dimension"] is None:
1045            if update:
1046                id = collection.id
1047                self._sysdb.update_collection(id=id, dimension=dim)
1048                collection["dimension"] = dim
1049        elif collection["dimension"] != dim:
1050            raise InvalidDimensionException(
1051                f"Embedding dimension {dim} does not match collection dimensionality {collection['dimension']}"
1052            )
1053        else:
1054            return  # all is well
1055
1056    @trace_method("SegmentAPI._get_collection", OpenTelemetryGranularity.ALL)
1057    def _get_collection(self, collection_id: UUID) -> t.Collection:
1058        collections = self._sysdb.get_collections(id=collection_id)
1059        if not collections or len(collections) == 0:
1060            raise NotFoundError(f"Collection {collection_id} does not exist.")
1061        return collections[0]
1062
1063    @trace_method("SegmentAPI._scan", OpenTelemetryGranularity.OPERATION)
1064    def _scan(self, collection_id: UUID) -> Scan:
1065        collection_and_segments = self._sysdb.get_collection_with_segments(
1066            collection_id
1067        )
1068        # For now collection should have exactly one segment per scope:
1069        # - Local scopes: vector, metadata
1070        # - Distributed scopes: vector, metadata, record
1071        scope_to_segment = {
1072            segment["scope"]: segment for segment in collection_and_segments["segments"]
1073        }
1074        return Scan(
1075            collection=collection_and_segments["collection"],
1076            knn=scope_to_segment[t.SegmentScope.VECTOR],
1077            metadata=scope_to_segment[t.SegmentScope.METADATA],
1078            # Local chroma do not have record segment, and this is not used by the local executor
1079            record=scope_to_segment.get(t.SegmentScope.RECORD, None),  # type: ignore[arg-type]
1080        )
1081
1082
1083def _records(
1084    operation: t.Operation,
1085    ids: IDs,
1086    embeddings: Optional[Embeddings] = None,
1087    metadatas: Optional[Metadatas] = None,
1088    documents: Optional[Documents] = None,
1089    uris: Optional[URIs] = None,
1090) -> Generator[t.OperationRecord, None, None]:
1091    """Convert parallel lists of embeddings, metadatas and documents to a sequence of
1092    SubmitEmbeddingRecords"""
1093
1094    # Presumes that callers were invoked via  Collection model, which means
1095    # that we know that the embeddings, metadatas and documents have already been
1096    # normalized and are guaranteed to be consistently named lists.
1097
1098    if embeddings == []:
1099        embeddings = None
1100
1101    for i, id in enumerate(ids):
1102        metadata = None
1103        if metadatas:
1104            metadata = metadatas[i]
1105
1106        if documents:
1107            document = documents[i]
1108            if metadata:
1109                metadata = {**metadata, "chroma:document": document}
1110            else:
1111                metadata = {"chroma:document": document}
1112
1113        if uris:
1114            uri = uris[i]
1115            if metadata:
1116                metadata = {**metadata, "chroma:uri": uri}
1117            else:
1118                metadata = {"chroma:uri": uri}
1119
1120        record = t.OperationRecord(
1121            id=id,
1122            embedding=embeddings[i] if embeddings is not None else None,
1123            encoding=t.ScalarEncoding.FLOAT32,  # Hardcode for now
1124            metadata=metadata,
1125            operation=operation,
1126        )
1127        yield record
1128 
codekingpro/portable-devtools · Team Ai