codekingpro/portable-devtools
114k
1from typing import Optional, Sequence
2from types import TracebackType
3from uuid import UUID
4
5from overrides import override
6import httpx
7from chromadb.api import AdminAPI, ClientAPI, ServerAPI
8from chromadb.api.collection_configuration import (
9 CreateCollectionConfiguration,
10 UpdateCollectionConfiguration,
11 validate_embedding_function_conflict_on_create,
12 validate_embedding_function_conflict_on_get,
13)
14from chromadb.api.shared_system_client import SharedSystemClient
15from chromadb.api.types import (
16 CollectionMetadata,
17 DataLoader,
18 Documents,
19 Embeddable,
20 EmbeddingFunction,
21 Embeddings,
22 GetResult,
23 IDs,
24 Include,
25 Loadable,
26 Metadatas,
27 QueryResult,
28 Schema,
29 URIs,
30 IncludeMetadataDocuments,
31 IncludeMetadataDocumentsDistances,
32 DefaultEmbeddingFunction,
33 DeleteResult,
34)
35from chromadb.auth import UserIdentity
36from chromadb.auth.utils import maybe_set_tenant_and_database
37from chromadb.config import Settings, System
38from chromadb.config import DEFAULT_TENANT, DEFAULT_DATABASE
39from chromadb.api.models.Collection import Collection
40from chromadb.errors import ChromaAuthError, ChromaError
41from chromadb.types import Database, Tenant, Where, WhereDocument
42
43
44class Client(SharedSystemClient, ClientAPI):
45 """A client for Chroma. This is the main entrypoint for interacting with Chroma.
46 A client internally stores its tenant and database and proxies calls to a
47 Server API instance of Chroma. It treats the Server API and corresponding System
48 as a singleton, so multiple clients connecting to the same resource will share the
49 same API instance.
50
51 Client implementations should be implement their own API-caching strategies.
52 """
53
54 tenant: str = DEFAULT_TENANT
55 database: str = DEFAULT_DATABASE
56
57 _server: ServerAPI
58 # An internal admin client for verifying that databases and tenants exist
59 _admin_client: AdminAPI
60 _closed: bool = False
61
62 # region Initialization
63 def __init__(
64 self,
65 tenant: Optional[str] = DEFAULT_TENANT,
66 database: Optional[str] = DEFAULT_DATABASE,
67 settings: Settings = Settings(),
68 ) -> None:
69 super().__init__(settings=settings)
70 try:
71 if tenant is not None:
72 self.tenant = tenant
73 if database is not None:
74 self.database = database
75
76 # Get the root system component we want to interact with
77 self._server = self._system.instance(ServerAPI)
78
79 user_identity = self.get_user_identity()
80
81 maybe_tenant, maybe_database = maybe_set_tenant_and_database(
82 user_identity,
83 overwrite_singleton_tenant_database_access_from_auth=settings.chroma_overwrite_singleton_tenant_database_access_from_auth,
84 user_provided_tenant=tenant,
85 user_provided_database=database,
86 )
87
88 # this should not happen unless types are invalidated
89 if maybe_tenant is None and tenant is None:
90 raise ChromaAuthError(
91 "Could not determine a tenant from the current authentication method. Please provide a tenant."
92 )
93 if maybe_database is None and database is None:
94 raise ChromaAuthError(
95 "Could not determine a database name from the current authentication method. Please provide a database name."
96 )
97
98 if maybe_tenant:
99 self.tenant = maybe_tenant
100 if maybe_database:
101 self.database = maybe_database
102
103 # Create an admin client for verifying that databases and tenants exist
104 self._admin_client = AdminClient.from_system(self._system)
105 self._validate_tenant_database(tenant=self.tenant, database=self.database)
106
107 self._submit_client_start_event()
108 except Exception:
109 # If init fails after refcount was incremented, release references
110 # to avoid a resource leak (the caller never receives the object to
111 # call close() on it).
112 if hasattr(self, "_admin_client"):
113 SharedSystemClient._release_system(self._admin_client._identifier)
114 SharedSystemClient._release_system(self._identifier)
115 raise
116
117 @classmethod
118 @override
119 def from_system(
120 cls,
121 system: System,
122 tenant: str = DEFAULT_TENANT,
123 database: str = DEFAULT_DATABASE,
124 ) -> "Client":
125 SharedSystemClient._populate_data_from_system(system)
126 instance = cls(tenant=tenant, database=database, settings=system.settings)
127 return instance
128
129 # endregion
130
131 @override
132 def get_user_identity(self) -> UserIdentity:
133 try:
134 return self._server.get_user_identity()
135 except httpx.ConnectError:
136 raise ValueError(
137 "Could not connect to a Chroma server. Are you sure it is running?"
138 )
139 # Propagate ChromaErrors
140 except ChromaError as e:
141 raise e
142 except Exception as e:
143 raise ValueError(str(e))
144
145 # region BaseAPI Methods
146 # Note - we could do this in less verbose ways, but they break type checking
147 @override
148 def heartbeat(self) -> int:
149 """Return the server time in nanoseconds since epoch."""
150 return self._server.heartbeat()
151
152 @override
153 def list_collections(
154 self, limit: Optional[int] = None, offset: Optional[int] = None
155 ) -> Sequence[Collection]:
156 """List collections for the current tenant and database, with pagination.
157
158 Returns:
159 Sequence[Collection]: Collection objects for the current tenant.
160 """
161 return [
162 Collection(client=self._server, model=model)
163 for model in self._server.list_collections(
164 limit, offset, tenant=self.tenant, database=self.database
165 )
166 ]
167
168 @override
169 def count_collections(self) -> int:
170 """Return the number of collections in the current database."""
171 return self._server.count_collections(
172 tenant=self.tenant, database=self.database
173 )
174
175 @override
176 def create_collection(
177 self,
178 name: str,
179 schema: Optional[Schema] = None,
180 configuration: Optional[CreateCollectionConfiguration] = None,
181 metadata: Optional[CollectionMetadata] = None,
182 embedding_function: Optional[
183 EmbeddingFunction[Embeddable]
184 ] = DefaultEmbeddingFunction(), # type: ignore
185 data_loader: Optional[DataLoader[Loadable]] = None,
186 get_or_create: bool = False,
187 ) -> Collection:
188 """Create a collection with optional configuration and metadata.
189
190 If using a schema, do not provide `embedding_function`. Instead,
191 provide the `embedding_function` as part of the schema.
192
193 Args:
194 name: Collection name.
195 schema: Optional collection schema for indexes and encryption.
196 configuration: Optional collection configuration.
197 metadata: Optional collection metadata.
198 embedding_function: Optional embedding function for the collection.
199 data_loader: Optional data loader for documents with URIs.
200 get_or_create: Whether to return an existing collection if present.
201
202 Returns:
203 Collection: The created collection.
204
205 Raises:
206 ValueError: If the embedding function conflicts with configuration.
207 """
208 if configuration is None:
209 configuration = {}
210
211 configuration_ef = configuration.get("embedding_function")
212
213 validate_embedding_function_conflict_on_create(
214 embedding_function, configuration_ef
215 )
216
217 # If ef provided in function params and collection config ef is None,
218 # set the collection config ef to the function params
219 if embedding_function is not None and configuration_ef is None:
220 configuration["embedding_function"] = embedding_function
221
222 model = self._server.create_collection(
223 name=name,
224 schema=schema,
225 metadata=metadata,
226 tenant=self.tenant,
227 database=self.database,
228 get_or_create=get_or_create,
229 configuration=configuration,
230 )
231 return Collection(
232 client=self._server,
233 model=model,
234 embedding_function=embedding_function,
235 data_loader=data_loader,
236 )
237
238 @override
239 def get_collection(
240 self,
241 name: str,
242 embedding_function: Optional[
243 EmbeddingFunction[Embeddable]
244 ] = DefaultEmbeddingFunction(), # type: ignore
245 data_loader: Optional[DataLoader[Loadable]] = None,
246 ) -> Collection:
247 """Get a collection by name.
248
249 Args:
250 name: Collection name.
251 embedding_function: Optional embedding function for the collection.
252 data_loader: Optional data loader for documents with URIs.
253
254 Returns:
255 Collection: The requested collection.
256
257 Raises:
258 ValueError: If the embedding function conflicts with configuration.
259 """
260 model = self._server.get_collection(
261 name=name,
262 tenant=self.tenant,
263 database=self.database,
264 )
265 persisted_ef_config = model.configuration_json.get("embedding_function")
266
267 validate_embedding_function_conflict_on_get(
268 embedding_function, persisted_ef_config
269 )
270
271 return Collection(
272 client=self._server,
273 model=model,
274 embedding_function=embedding_function,
275 data_loader=data_loader,
276 )
277
278 @override
279 def get_collection_by_id(
280 self,
281 id: UUID,
282 embedding_function: Optional[
283 EmbeddingFunction[Embeddable]
284 ] = DefaultEmbeddingFunction(), # type: ignore
285 data_loader: Optional[DataLoader[Loadable]] = None,
286 ) -> Collection:
287 """Get a collection by its ID.
288
289 Args:
290 id: The UUID of the collection.
291 embedding_function: Optional embedding function for the collection.
292 data_loader: Optional data loader for documents with URIs.
293
294 Returns:
295 Collection: The requested collection.
296
297 Raises:
298 ValueError: If the embedding function conflicts with configuration.
299 """
300 model = self._server.get_collection_by_id(
301 collection_id=id,
302 tenant=self.tenant,
303 database=self.database,
304 )
305 persisted_ef_config = model.configuration_json.get("embedding_function")
306
307 validate_embedding_function_conflict_on_get(
308 embedding_function, persisted_ef_config
309 )
310
311 return Collection(
312 client=self._server,
313 model=model,
314 embedding_function=embedding_function,
315 data_loader=data_loader,
316 )
317
318 @override
319 def get_or_create_collection(
320 self,
321 name: str,
322 schema: Optional[Schema] = None,
323 configuration: Optional[CreateCollectionConfiguration] = None,
324 metadata: Optional[CollectionMetadata] = None,
325 embedding_function: Optional[
326 EmbeddingFunction[Embeddable]
327 ] = DefaultEmbeddingFunction(), # type: ignore
328 data_loader: Optional[DataLoader[Loadable]] = None,
329 ) -> Collection:
330 """Get an existing collection or create a new one.
331
332 If the collection does not exist, it will be created. If the collection
333 already exists, the schema, configuration, and metadata arguments
334 will be ignored.
335
336 Args:
337 name: Collection name.
338 schema: Optional collection schema for indexes and encryption.
339 configuration: Optional collection configuration.
340 metadata: Optional collection metadata.
341 embedding_function: Optional embedding function for the collection.
342 data_loader: Optional data loader for URI-backed data.
343
344 Returns:
345 Collection: The existing or newly created collection.
346
347 Raises:
348 ValueError: If the embedding function does not match the collection's embedding function.
349 """
350 if configuration is None:
351 configuration = {}
352
353 configuration_ef = configuration.get("embedding_function")
354
355 validate_embedding_function_conflict_on_create(
356 embedding_function, configuration_ef
357 )
358
359 if embedding_function is not None and configuration_ef is None:
360 configuration["embedding_function"] = embedding_function
361 model = self._server.get_or_create_collection(
362 name=name,
363 schema=schema,
364 metadata=metadata,
365 tenant=self.tenant,
366 database=self.database,
367 configuration=configuration,
368 )
369
370 persisted_ef_config = model.configuration_json.get("embedding_function")
371
372 validate_embedding_function_conflict_on_get(
373 embedding_function, persisted_ef_config
374 )
375
376 return Collection(
377 client=self._server,
378 model=model,
379 embedding_function=embedding_function,
380 data_loader=data_loader,
381 )
382
383 @override
384 def _modify(
385 self,
386 id: UUID,
387 new_name: Optional[str] = None,
388 new_metadata: Optional[CollectionMetadata] = None,
389 new_configuration: Optional[UpdateCollectionConfiguration] = None,
390 ) -> None:
391 return self._server._modify(
392 id=id,
393 tenant=self.tenant,
394 database=self.database,
395 new_name=new_name,
396 new_metadata=new_metadata,
397 new_configuration=new_configuration,
398 )
399
400 @override
401 def delete_collection(
402 self,
403 name: str,
404 ) -> None:
405 return self._server.delete_collection(
406 name=name,
407 tenant=self.tenant,
408 database=self.database,
409 )
410
411 #
412 # ITEM METHODS
413 #
414
415 @override
416 def _add(
417 self,
418 ids: IDs,
419 collection_id: UUID,
420 embeddings: Embeddings,
421 metadatas: Optional[Metadatas] = None,
422 documents: Optional[Documents] = None,
423 uris: Optional[URIs] = None,
424 ) -> bool:
425 return self._server._add(
426 ids=ids,
427 tenant=self.tenant,
428 database=self.database,
429 collection_id=collection_id,
430 embeddings=embeddings,
431 metadatas=metadatas,
432 documents=documents,
433 uris=uris,
434 )
435
436 @override
437 def _update(
438 self,
439 collection_id: UUID,
440 ids: IDs,
441 embeddings: Optional[Embeddings] = None,
442 metadatas: Optional[Metadatas] = None,
443 documents: Optional[Documents] = None,
444 uris: Optional[URIs] = None,
445 ) -> bool:
446 return self._server._update(
447 collection_id=collection_id,
448 tenant=self.tenant,
449 database=self.database,
450 ids=ids,
451 embeddings=embeddings,
452 metadatas=metadatas,
453 documents=documents,
454 uris=uris,
455 )
456
457 @override
458 def _upsert(
459 self,
460 collection_id: UUID,
461 ids: IDs,
462 embeddings: Embeddings,
463 metadatas: Optional[Metadatas] = None,
464 documents: Optional[Documents] = None,
465 uris: Optional[URIs] = None,
466 ) -> bool:
467 return self._server._upsert(
468 collection_id=collection_id,
469 tenant=self.tenant,
470 database=self.database,
471 ids=ids,
472 embeddings=embeddings,
473 metadatas=metadatas,
474 documents=documents,
475 uris=uris,
476 )
477
478 @override
479 def _count(self, collection_id: UUID) -> int:
480 return self._server._count(
481 collection_id=collection_id,
482 tenant=self.tenant,
483 database=self.database,
484 )
485
486 @override
487 def _peek(self, collection_id: UUID, n: int = 10) -> GetResult:
488 return self._server._peek(
489 collection_id=collection_id,
490 n=n,
491 tenant=self.tenant,
492 database=self.database,
493 )
494
495 @override
496 def _get(
497 self,
498 collection_id: UUID,
499 ids: Optional[IDs] = None,
500 where: Optional[Where] = None,
501 limit: Optional[int] = None,
502 offset: Optional[int] = None,
503 where_document: Optional[WhereDocument] = None,
504 include: Include = IncludeMetadataDocuments,
505 ) -> GetResult:
506 return self._server._get(
507 collection_id=collection_id,
508 tenant=self.tenant,
509 database=self.database,
510 ids=ids,
511 where=where,
512 limit=limit,
513 offset=offset,
514 where_document=where_document,
515 include=include,
516 )
517
518 def _delete(
519 self,
520 collection_id: UUID,
521 ids: Optional[IDs],
522 where: Optional[Where] = None,
523 where_document: Optional[WhereDocument] = None,
524 limit: Optional[int] = None,
525 ) -> DeleteResult:
526 return self._server._delete(
527 collection_id=collection_id,
528 tenant=self.tenant,
529 database=self.database,
530 ids=ids,
531 where=where,
532 where_document=where_document,
533 limit=limit,
534 )
535
536 @override
537 def _query(
538 self,
539 collection_id: UUID,
540 query_embeddings: Embeddings,
541 ids: Optional[IDs] = None,
542 n_results: int = 10,
543 where: Optional[Where] = None,
544 where_document: Optional[WhereDocument] = None,
545 include: Include = IncludeMetadataDocumentsDistances,
546 ) -> QueryResult:
547 return self._server._query(
548 collection_id=collection_id,
549 ids=ids,
550 tenant=self.tenant,
551 database=self.database,
552 query_embeddings=query_embeddings,
553 n_results=n_results,
554 where=where,
555 where_document=where_document,
556 include=include,
557 )
558
559 @override
560 def reset(self) -> bool:
561 return self._server.reset()
562
563 @override
564 def get_version(self) -> str:
565 return self._server.get_version()
566
567 @override
568 def get_settings(self) -> Settings:
569 return self._server.get_settings()
570
571 @override
572 def get_max_batch_size(self) -> int:
573 return self._server.get_max_batch_size()
574
575 # endregion
576
577 # region ClientAPI Methods
578
579 @override
580 def set_tenant(self, tenant: str, database: str = DEFAULT_DATABASE) -> None:
581 self._validate_tenant_database(tenant=tenant, database=database)
582 self.tenant = tenant
583 self.database = database
584
585 @override
586 def set_database(self, database: str) -> None:
587 self._validate_tenant_database(tenant=self.tenant, database=database)
588 self.database = database
589
590 def close(self) -> None:
591 """Close the client and release all resources.
592
593 This method decrements the reference count for the underlying System.
594 When the last client using a shared System calls close(), the System
595 is stopped and all resources (database connections, etc.) are released.
596
597 This is particularly important for PersistentClient to avoid SQLite
598 file locking issues.
599
600 Note: If multiple clients share the same System (e.g., multiple PersistentClient
601 instances with the same path), the System will only be stopped when the last
602 client is closed. This allows safe use of context managers with multiple clients.
603
604 Example:
605 >>> client = chromadb.PersistentClient(path="./chroma_db")
606 >>> # ... use client ...
607 >>> client.close()
608
609 Or using context manager:
610 >>> with chromadb.PersistentClient(path="./chroma_db") as client:
611 ... # ... use client ...
612 """
613 # Make close() idempotent - a second call is a safe no-op
614 if self._closed:
615 return
616 self._closed = True
617
618 # Release the internal admin client's reference first, since it also
619 # incremented the refcount for the shared system on creation.
620 if hasattr(self, "_admin_client"):
621 SharedSystemClient._release_system(self._admin_client._identifier)
622
623 # Release our own reference; stops system if this was the last client
624 SharedSystemClient._release_system(self._identifier)
625
626 def __enter__(self) -> "Client":
627 """Context manager entry."""
628 return self
629
630 def __exit__(
631 self,
632 exc_type: Optional[type[BaseException]],
633 exc_val: Optional[BaseException],
634 exc_tb: Optional[TracebackType],
635 ) -> None:
636 """Context manager exit."""
637 self.close()
638
639 def _validate_tenant_database(self, tenant: str, database: str) -> None:
640 try:
641 self._admin_client.get_tenant(name=tenant)
642 except httpx.ConnectError:
643 raise ValueError(
644 "Could not connect to a Chroma server. Are you sure it is running?"
645 )
646 # Propagate ChromaErrors
647 except ChromaError as e:
648 raise e
649 except Exception:
650 raise ValueError(
651 f"Could not connect to tenant {tenant}. Are you sure it exists?"
652 )
653
654 try:
655 self._admin_client.get_database(name=database, tenant=tenant)
656 except httpx.ConnectError:
657 raise ValueError(
658 "Could not connect to a Chroma server. Are you sure it is running?"
659 )
660
661 # endregion
662
663
664class AdminClient(SharedSystemClient, AdminAPI):
665 """Admin client for managing tenants and databases."""
666
667 _server: ServerAPI
668
669 def __init__(self, settings: Settings = Settings()) -> None:
670 super().__init__(settings)
671 self._server = self._system.instance(ServerAPI)
672
673 @override
674 def create_database(self, name: str, tenant: str = DEFAULT_TENANT) -> None:
675 """Create a database in a tenant.
676
677 Args:
678 name: Database name.
679 tenant: Tenant that owns the database.
680 """
681 return self._server.create_database(name=name, tenant=tenant)
682
683 @override
684 def get_database(self, name: str, tenant: str = DEFAULT_TENANT) -> Database:
685 """Get a database by name.
686
687 Args:
688 name: Database name.
689 tenant: Tenant that owns the database.
690
691 Returns:
692 Database: The database record.
693 """
694 return self._server.get_database(name=name, tenant=tenant)
695
696 @override
697 def delete_database(self, name: str, tenant: str = DEFAULT_TENANT) -> None:
698 """Delete a database by name.
699
700 Args:
701 name: Database name.
702 tenant: Tenant that owns the database.
703 """
704 return self._server.delete_database(name=name, tenant=tenant)
705
706 @override
707 def list_databases(
708 self,
709 limit: Optional[int] = None,
710 offset: Optional[int] = None,
711 tenant: str = DEFAULT_TENANT,
712 ) -> Sequence[Database]:
713 return self._server.list_databases(limit, offset, tenant=tenant)
714
715 @override
716 def create_tenant(self, name: str) -> None:
717 return self._server.create_tenant(name=name)
718
719 @override
720 def get_tenant(self, name: str) -> Tenant:
721 return self._server.get_tenant(name=name)
722
723 @classmethod
724 @override
725 def from_system(
726 cls,
727 system: System,
728 ) -> "AdminClient":
729 SharedSystemClient._populate_data_from_system(system)
730 instance = cls(settings=system.settings)
731 return instance
732 