Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
client.py732 linesDownload Raw Back to api
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 
codekingpro/portable-devtools · Team Ai