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