Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
async_client.py570 linesDownload Raw Back to api
1import httpx
2from typing import Optional, Sequence
3from uuid import UUID
4from overrides import override
5
6from chromadb.auth import UserIdentity
7from chromadb.auth.utils import maybe_set_tenant_and_database
8from chromadb.api import AsyncAdminAPI, AsyncClientAPI, AsyncServerAPI
9from chromadb.api.collection_configuration import (
10    CreateCollectionConfiguration,
11    UpdateCollectionConfiguration,
12    validate_embedding_function_conflict_on_create,
13    validate_embedding_function_conflict_on_get,
14)
15from chromadb.api.models.AsyncCollection import AsyncCollection
16from chromadb.api.shared_system_client import SharedSystemClient
17from chromadb.api.types import (
18    CollectionMetadata,
19    DataLoader,
20    Documents,
21    Embeddable,
22    EmbeddingFunction,
23    Embeddings,
24    GetResult,
25    IDs,
26    Include,
27    IncludeMetadataDocuments,
28    IncludeMetadataDocumentsDistances,
29    Loadable,
30    Metadatas,
31    QueryResult,
32    Schema,
33    URIs,
34    DefaultEmbeddingFunction,
35    DeleteResult,
36)
37from chromadb.config import DEFAULT_DATABASE, DEFAULT_TENANT, Settings, System
38from chromadb.errors import ChromaError
39from chromadb.types import Database, Tenant, Where, WhereDocument
40
41
42class AsyncClient(SharedSystemClient, AsyncClientAPI):
43    """A client for Chroma. This is the main entrypoint for interacting with Chroma.
44    A client internally stores its tenant and database and proxies calls to a
45    Server API instance of Chroma. It treats the Server API and corresponding System
46    as a singleton, so multiple clients connecting to the same resource will share the
47    same API instance.
48
49    Client implementations should be implement their own API-caching strategies.
50    """
51
52    # An internal admin client for verifying that databases and tenants exist
53    _admin_client: AsyncAdminAPI
54
55    tenant: str = DEFAULT_TENANT
56    database: str = DEFAULT_DATABASE
57
58    _server: AsyncServerAPI
59
60    @classmethod
61    async def create(
62        cls,
63        tenant: str = DEFAULT_TENANT,
64        database: str = DEFAULT_DATABASE,
65        settings: Settings = Settings(),
66    ) -> "AsyncClient":
67        # Create an admin client for verifying that databases and tenants exist
68        self = cls(settings=settings)
69        SharedSystemClient._populate_data_from_system(self._system)
70
71        self.tenant = tenant
72        self.database = database
73
74        # Get the root system component we want to interact with
75        self._server = self._system.instance(AsyncServerAPI)
76
77        user_identity = await self.get_user_identity()
78
79        maybe_tenant, maybe_database = maybe_set_tenant_and_database(
80            user_identity,
81            overwrite_singleton_tenant_database_access_from_auth=settings.chroma_overwrite_singleton_tenant_database_access_from_auth,
82            user_provided_tenant=tenant,
83            user_provided_database=database,
84        )
85        if maybe_tenant:
86            self.tenant = maybe_tenant
87        if maybe_database:
88            self.database = maybe_database
89
90        self._admin_client = AsyncAdminClient.from_system(self._system)
91        await self._validate_tenant_database(tenant=self.tenant, database=self.database)
92
93        self._submit_client_start_event()
94
95        return self
96
97    @classmethod
98    # (we can't override and use from_system() because it's synchronous)
99    async def from_system_async(
100        cls,
101        system: System,
102        tenant: str = DEFAULT_TENANT,
103        database: str = DEFAULT_DATABASE,
104    ) -> "AsyncClient":
105        """Create a client from an existing system. This is useful for testing and debugging."""
106        return await AsyncClient.create(tenant, database, system.settings)
107
108    @classmethod
109    @override
110    def from_system(
111        cls,
112        system: System,
113    ) -> "SharedSystemClient":
114        """AsyncClient cannot be created synchronously. Use .from_system_async() instead."""
115        raise NotImplementedError(
116            "AsyncClient cannot be created synchronously. Use .from_system_async() instead."
117        )
118
119    @override
120    async def get_user_identity(self) -> UserIdentity:
121        return await self._server.get_user_identity()
122
123    @override
124    async def set_tenant(self, tenant: str, database: str = DEFAULT_DATABASE) -> None:
125        await self._validate_tenant_database(tenant=tenant, database=database)
126        self.tenant = tenant
127        self.database = database
128
129    @override
130    async def set_database(self, database: str) -> None:
131        await self._validate_tenant_database(tenant=self.tenant, database=database)
132        self.database = database
133
134    async def _validate_tenant_database(self, tenant: str, database: str) -> None:
135        try:
136            await self._admin_client.get_tenant(name=tenant)
137        except httpx.ConnectError:
138            raise ValueError(
139                "Could not connect to a Chroma server. Are you sure it is running?"
140            )
141        # Propagate ChromaErrors
142        except ChromaError as e:
143            raise e
144        except Exception:
145            raise ValueError(
146                f"Could not connect to tenant {tenant}. Are you sure it exists?"
147            )
148
149        try:
150            await self._admin_client.get_database(name=database, tenant=tenant)
151        except httpx.ConnectError:
152            raise ValueError(
153                "Could not connect to a Chroma server. Are you sure it is running?"
154            )
155
156    # region BaseAPI Methods
157    # Note - we could do this in less verbose ways, but they break type checking
158    @override
159    async def heartbeat(self) -> int:
160        return await self._server.heartbeat()
161
162    @override
163    async def list_collections(
164        self, limit: Optional[int] = None, offset: Optional[int] = None
165    ) -> Sequence[AsyncCollection]:
166        models = await self._server.list_collections(
167            limit, offset, tenant=self.tenant, database=self.database
168        )
169        return [AsyncCollection(client=self._server, model=model) for model in models]
170
171    @override
172    async def count_collections(self) -> int:
173        return await self._server.count_collections(
174            tenant=self.tenant, database=self.database
175        )
176
177    @override
178    async def create_collection(
179        self,
180        name: str,
181        schema: Optional[Schema] = None,
182        configuration: Optional[CreateCollectionConfiguration] = None,
183        metadata: Optional[CollectionMetadata] = None,
184        embedding_function: Optional[
185            EmbeddingFunction[Embeddable]
186        ] = DefaultEmbeddingFunction(),  # type: ignore
187        data_loader: Optional[DataLoader[Loadable]] = None,
188        get_or_create: bool = False,
189    ) -> AsyncCollection:
190        if configuration is None:
191            configuration = {}
192
193        configuration_ef = configuration.get("embedding_function")
194
195        validate_embedding_function_conflict_on_create(
196            embedding_function, configuration_ef
197        )
198
199        # If ef provided in function params and collection config ef is None,
200        # set the collection config ef to the function params
201        if embedding_function is not None and configuration_ef is None:
202            configuration["embedding_function"] = embedding_function
203
204        model = await self._server.create_collection(
205            name=name,
206            schema=schema,
207            configuration=configuration,
208            metadata=metadata,
209            tenant=self.tenant,
210            database=self.database,
211            get_or_create=get_or_create,
212        )
213        return AsyncCollection(
214            client=self._server,
215            model=model,
216            embedding_function=embedding_function,
217            data_loader=data_loader,
218        )
219
220    @override
221    async def get_collection(
222        self,
223        name: str,
224        embedding_function: Optional[
225            EmbeddingFunction[Embeddable]
226        ] = DefaultEmbeddingFunction(),  # type: ignore
227        data_loader: Optional[DataLoader[Loadable]] = None,
228    ) -> AsyncCollection:
229        model = await self._server.get_collection(
230            name=name,
231            tenant=self.tenant,
232            database=self.database,
233        )
234        persisted_ef_config = model.configuration_json.get("embedding_function")
235
236        validate_embedding_function_conflict_on_get(
237            embedding_function, persisted_ef_config
238        )
239
240        return AsyncCollection(
241            client=self._server,
242            model=model,
243            embedding_function=embedding_function,
244            data_loader=data_loader,
245        )
246
247    @override
248    async def get_collection_by_id(
249        self,
250        id: UUID,
251        embedding_function: Optional[
252            EmbeddingFunction[Embeddable]
253        ] = DefaultEmbeddingFunction(),  # type: ignore
254        data_loader: Optional[DataLoader[Loadable]] = None,
255    ) -> AsyncCollection:
256        """Get a collection by its ID.
257
258        Args:
259            id: The UUID of the collection.
260            embedding_function: Optional embedding function for the collection.
261            data_loader: Optional data loader for documents with URIs.
262
263        Returns:
264            AsyncCollection: The requested collection.
265
266        Raises:
267            ValueError: If the embedding function conflicts with configuration.
268        """
269        model = await self._server.get_collection_by_id(
270            collection_id=id,
271            tenant=self.tenant,
272            database=self.database,
273        )
274        persisted_ef_config = model.configuration_json.get("embedding_function")
275
276        validate_embedding_function_conflict_on_get(
277            embedding_function, persisted_ef_config
278        )
279
280        return AsyncCollection(
281            client=self._server,
282            model=model,
283            embedding_function=embedding_function,
284            data_loader=data_loader,
285        )
286
287    @override
288    async def get_or_create_collection(
289        self,
290        name: str,
291        schema: Optional[Schema] = None,
292        configuration: Optional[CreateCollectionConfiguration] = None,
293        metadata: Optional[CollectionMetadata] = None,
294        embedding_function: Optional[
295            EmbeddingFunction[Embeddable]
296        ] = DefaultEmbeddingFunction(),  # type: ignore
297        data_loader: Optional[DataLoader[Loadable]] = None,
298    ) -> AsyncCollection:
299        if configuration is None:
300            configuration = {}
301
302        configuration_ef = configuration.get("embedding_function")
303
304        validate_embedding_function_conflict_on_create(
305            embedding_function, configuration_ef
306        )
307
308        if embedding_function is not None and configuration_ef is None:
309            configuration["embedding_function"] = embedding_function
310        model = await self._server.get_or_create_collection(
311            name=name,
312            schema=schema,
313            configuration=configuration,
314            metadata=metadata,
315            tenant=self.tenant,
316            database=self.database,
317        )
318
319        persisted_ef_config = model.configuration_json.get("embedding_function")
320
321        validate_embedding_function_conflict_on_get(
322            embedding_function, persisted_ef_config
323        )
324
325        return AsyncCollection(
326            client=self._server,
327            model=model,
328            embedding_function=embedding_function,
329            data_loader=data_loader,
330        )
331
332    @override
333    async def _modify(
334        self,
335        id: UUID,
336        new_name: Optional[str] = None,
337        new_metadata: Optional[CollectionMetadata] = None,
338        new_configuration: Optional[UpdateCollectionConfiguration] = None,
339    ) -> None:
340        return await self._server._modify(
341            id=id,
342            new_name=new_name,
343            new_metadata=new_metadata,
344            new_configuration=new_configuration,
345            tenant=self.tenant,
346            database=self.database,
347        )
348
349    @override
350    async def delete_collection(
351        self,
352        name: str,
353    ) -> None:
354        return await self._server.delete_collection(
355            name=name,
356            tenant=self.tenant,
357            database=self.database,
358        )
359
360    #
361    # ITEM METHODS
362    #
363
364    @override
365    async def _add(
366        self,
367        ids: IDs,
368        collection_id: UUID,
369        embeddings: Embeddings,
370        metadatas: Optional[Metadatas] = None,
371        documents: Optional[Documents] = None,
372        uris: Optional[URIs] = None,
373    ) -> bool:
374        return await self._server._add(
375            ids=ids,
376            collection_id=collection_id,
377            embeddings=embeddings,
378            metadatas=metadatas,
379            documents=documents,
380            uris=uris,
381            tenant=self.tenant,
382            database=self.database,
383        )
384
385    @override
386    async def _update(
387        self,
388        collection_id: UUID,
389        ids: IDs,
390        embeddings: Optional[Embeddings] = None,
391        metadatas: Optional[Metadatas] = None,
392        documents: Optional[Documents] = None,
393        uris: Optional[URIs] = None,
394    ) -> bool:
395        return await self._server._update(
396            collection_id=collection_id,
397            ids=ids,
398            embeddings=embeddings,
399            metadatas=metadatas,
400            documents=documents,
401            uris=uris,
402            tenant=self.tenant,
403            database=self.database,
404        )
405
406    @override
407    async def _upsert(
408        self,
409        collection_id: UUID,
410        ids: IDs,
411        embeddings: Embeddings,
412        metadatas: Optional[Metadatas] = None,
413        documents: Optional[Documents] = None,
414        uris: Optional[URIs] = None,
415    ) -> bool:
416        return await self._server._upsert(
417            collection_id=collection_id,
418            ids=ids,
419            embeddings=embeddings,
420            metadatas=metadatas,
421            documents=documents,
422            uris=uris,
423            tenant=self.tenant,
424            database=self.database,
425        )
426
427    @override
428    async def _count(self, collection_id: UUID) -> int:
429        return await self._server._count(
430            collection_id=collection_id,
431        )
432
433    @override
434    async def _peek(self, collection_id: UUID, n: int = 10) -> GetResult:
435        return await self._server._peek(
436            collection_id=collection_id,
437            n=n,
438        )
439
440    @override
441    async def _get(
442        self,
443        collection_id: UUID,
444        ids: Optional[IDs] = None,
445        where: Optional[Where] = None,
446        limit: Optional[int] = None,
447        offset: Optional[int] = None,
448        where_document: Optional[WhereDocument] = None,
449        include: Include = IncludeMetadataDocuments,
450    ) -> GetResult:
451        return await self._server._get(
452            collection_id=collection_id,
453            ids=ids,
454            where=where,
455            limit=limit,
456            offset=offset,
457            where_document=where_document,
458            include=include,
459            tenant=self.tenant,
460            database=self.database,
461        )
462
463    async def _delete(
464        self,
465        collection_id: UUID,
466        ids: Optional[IDs],
467        where: Optional[Where] = None,
468        where_document: Optional[WhereDocument] = None,
469        limit: Optional[int] = None,
470    ) -> DeleteResult:
471        return await self._server._delete(
472            collection_id=collection_id,
473            ids=ids,
474            where=where,
475            where_document=where_document,
476            limit=limit,
477            tenant=self.tenant,
478            database=self.database,
479        )
480
481    @override
482    async def _query(
483        self,
484        collection_id: UUID,
485        query_embeddings: Embeddings,
486        ids: Optional[IDs] = None,
487        n_results: int = 10,
488        where: Optional[Where] = None,
489        where_document: Optional[WhereDocument] = None,
490        include: Include = IncludeMetadataDocumentsDistances,
491    ) -> QueryResult:
492        return await self._server._query(
493            collection_id=collection_id,
494            query_embeddings=query_embeddings,
495            ids=ids,
496            n_results=n_results,
497            where=where,
498            where_document=where_document,
499            include=include,
500            tenant=self.tenant,
501            database=self.database,
502        )
503
504    @override
505    async def reset(self) -> bool:
506        return await self._server.reset()
507
508    @override
509    async def get_version(self) -> str:
510        return await self._server.get_version()
511
512    @override
513    def get_settings(self) -> Settings:
514        return self._server.get_settings()
515
516    @override
517    async def get_max_batch_size(self) -> int:
518        return await self._server.get_max_batch_size()
519
520    # endregion
521
522
523class AsyncAdminClient(SharedSystemClient, AsyncAdminAPI):
524    _server: AsyncServerAPI
525
526    def __init__(self, settings: Settings = Settings()) -> None:
527        super().__init__(settings)
528        self._server = self._system.instance(AsyncServerAPI)
529
530    @override
531    async def create_database(self, name: str, tenant: str = DEFAULT_TENANT) -> None:
532        return await self._server.create_database(name=name, tenant=tenant)
533
534    @override
535    async def get_database(self, name: str, tenant: str = DEFAULT_TENANT) -> Database:
536        return await self._server.get_database(name=name, tenant=tenant)
537
538    @override
539    async def delete_database(self, name: str, tenant: str = DEFAULT_TENANT) -> None:
540        return await self._server.delete_database(name=name, tenant=tenant)
541
542    @override
543    async def list_databases(
544        self,
545        limit: Optional[int] = None,
546        offset: Optional[int] = None,
547        tenant: str = DEFAULT_TENANT,
548    ) -> Sequence[Database]:
549        return await self._server.list_databases(
550            limit=limit, offset=offset, tenant=tenant
551        )
552
553    @override
554    async def create_tenant(self, name: str) -> None:
555        return await self._server.create_tenant(name=name)
556
557    @override
558    async def get_tenant(self, name: str) -> Tenant:
559        return await self._server.get_tenant(name=name)
560
561    @classmethod
562    @override
563    def from_system(
564        cls,
565        system: System,
566    ) -> "AsyncAdminClient":
567        SharedSystemClient._populate_data_from_system(system)
568        instance = cls(settings=system.settings)
569        return instance
570 
codekingpro/portable-devtools · Team Ai