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