codekingpro/portable-devtools
115k
1import asyncio
2from uuid import UUID
3import urllib.parse
4import orjson
5from typing import Any, Mapping, Optional, cast, Tuple, Sequence, Dict, List
6import logging
7import httpx
8from overrides import override
9from chromadb import __version__
10from chromadb.auth import UserIdentity
11from chromadb.api.async_api import AsyncServerAPI
12from chromadb.api.base_http_client import BaseHTTPClient
13from chromadb.api.collection_configuration import (
14 CreateCollectionConfiguration,
15 UpdateCollectionConfiguration,
16 create_collection_configuration_to_json,
17 update_collection_configuration_to_json,
18)
19from chromadb.config import DEFAULT_DATABASE, DEFAULT_TENANT, System, Settings
20from chromadb.telemetry.opentelemetry import (
21 OpenTelemetryClient,
22 OpenTelemetryGranularity,
23 trace_method,
24)
25from chromadb.telemetry.product import ProductTelemetryClient
26from chromadb.utils.async_to_sync import async_to_sync
27from chromadb.types import Database, Tenant, Collection as CollectionModel
28from chromadb.execution.expression.plan import Search
29
30from chromadb.api.types import (
31 DeleteResult,
32 Documents,
33 Embeddings,
34 IDs,
35 Include,
36 IndexingStatus,
37 Schema,
38 Metadatas,
39 ReadLevel,
40 URIs,
41 Where,
42 WhereDocument,
43 GetResult,
44 QueryResult,
45 SearchResult,
46 CollectionMetadata,
47 optional_embeddings_to_base64_strings,
48 validate_batch,
49 convert_np_embeddings_to_list,
50 IncludeMetadataDocuments,
51 IncludeMetadataDocumentsDistances,
52)
53
54from chromadb.api.types import (
55 IncludeMetadataDocumentsEmbeddings,
56 serialize_metadata,
57 deserialize_metadata,
58)
59
60
61logger = logging.getLogger(__name__)
62
63
64class AsyncFastAPI(BaseHTTPClient, AsyncServerAPI):
65 # We make one client per event loop to avoid unexpected issues if a client
66 # is shared between event loops.
67 # For example, if a client is constructed in the main thread, then passed
68 # (or a returned Collection is passed) to a new thread, the client would
69 # normally throw an obscure asyncio error.
70 # Mixing asyncio and threading in this manner usually discouraged, but
71 # this gives a better user experience with practically no downsides.
72 # https://github.com/encode/httpx/issues/2058
73 _clients: Dict[int, httpx.AsyncClient] = {}
74
75 def __init__(self, system: System):
76 super().__init__(system)
77
78 system.settings.require("chroma_server_host")
79 system.settings.require("chroma_server_http_port")
80
81 self._opentelemetry_client = self.require(OpenTelemetryClient)
82 self._product_telemetry_client = self.require(ProductTelemetryClient)
83 self._settings = system.settings
84
85 self._api_url = AsyncFastAPI.resolve_url(
86 chroma_server_host=str(system.settings.chroma_server_host),
87 chroma_server_http_port=system.settings.chroma_server_http_port,
88 chroma_server_ssl_enabled=system.settings.chroma_server_ssl_enabled,
89 default_api_path=system.settings.chroma_server_api_default_path,
90 )
91
92 async def __aenter__(self) -> "AsyncFastAPI":
93 self._get_client()
94 return self
95
96 async def _cleanup(self) -> None:
97 while len(self._clients) > 0:
98 (_, client) = self._clients.popitem()
99 await client.aclose()
100
101 async def __aexit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
102 await self._cleanup()
103
104 @override
105 def stop(self) -> None:
106 super().stop()
107
108 @async_to_sync
109 async def sync_cleanup() -> None:
110 await self._cleanup()
111
112 sync_cleanup()
113
114 def _get_client(self) -> httpx.AsyncClient:
115 # Ideally this would use anyio to be compatible with both
116 # asyncio and trio, but anyio does not expose any way to identify
117 # the current event loop.
118 # We attempt to get the loop assuming the environment is asyncio, and
119 # otherwise gracefully fall back to using a singleton client.
120 loop_hash = None
121 try:
122 loop = asyncio.get_event_loop()
123 loop_hash = loop.__hash__()
124 except RuntimeError:
125 loop_hash = 0
126
127 if loop_hash not in self._clients:
128 headers = (self._settings.chroma_server_headers or {}).copy()
129 headers["Content-Type"] = "application/json"
130 headers["User-Agent"] = (
131 "Chroma Python Client v"
132 + __version__
133 + " (https://github.com/chroma-core/chroma)"
134 )
135
136 self._clients[loop_hash] = httpx.AsyncClient(
137 timeout=None,
138 headers=headers,
139 verify=self._settings.chroma_server_ssl_verify or False,
140 limits=self.http_limits,
141 )
142
143 return self._clients[loop_hash]
144
145 @override
146 def get_request_headers(self) -> Mapping[str, str]:
147 return dict(self._get_client().headers)
148
149 @override
150 def get_api_url(self) -> str:
151 return self._api_url
152
153 async def _make_request(
154 self, method: str, path: str, **kwargs: Dict[str, Any]
155 ) -> Any:
156 # If the request has json in kwargs, use orjson to serialize it,
157 # remove it from kwargs, and add it to the content parameter
158 # This is because httpx uses a slower json serializer
159 if "json" in kwargs:
160 data = orjson.dumps(kwargs.pop("json"), option=orjson.OPT_SERIALIZE_NUMPY)
161 kwargs["content"] = data
162
163 # Unlike requests, httpx does not automatically escape the path
164 escaped_path = urllib.parse.quote(path, safe="/", encoding=None, errors=None)
165 url = self._api_url + escaped_path
166
167 response = await self._get_client().request(method, url, **cast(Any, kwargs))
168 BaseHTTPClient._raise_chroma_error(response)
169 return orjson.loads(response.text)
170
171 @trace_method("AsyncFastAPI.heartbeat", OpenTelemetryGranularity.OPERATION)
172 @override
173 async def heartbeat(self) -> int:
174 response = await self._make_request("get", "")
175 return int(response["nanosecond heartbeat"])
176
177 @trace_method("AsyncFastAPI.create_database", OpenTelemetryGranularity.OPERATION)
178 @override
179 async def create_database(
180 self,
181 name: str,
182 tenant: str = DEFAULT_TENANT,
183 ) -> None:
184 await self._make_request(
185 "post",
186 f"/tenants/{tenant}/databases",
187 json={"name": name},
188 )
189
190 @trace_method("AsyncFastAPI.get_database", OpenTelemetryGranularity.OPERATION)
191 @override
192 async def get_database(
193 self,
194 name: str,
195 tenant: str = DEFAULT_TENANT,
196 ) -> Database:
197 response = await self._make_request(
198 "get",
199 f"/tenants/{tenant}/databases/{name}",
200 params={"tenant": tenant},
201 )
202
203 return Database(
204 id=response["id"], name=response["name"], tenant=response["tenant"]
205 )
206
207 @trace_method("AsyncFastAPI.delete_database", OpenTelemetryGranularity.OPERATION)
208 @override
209 async def delete_database(
210 self,
211 name: str,
212 tenant: str = DEFAULT_TENANT,
213 ) -> None:
214 await self._make_request(
215 "delete",
216 f"/tenants/{tenant}/databases/{name}",
217 )
218
219 @trace_method("AsyncFastAPI.list_databases", OpenTelemetryGranularity.OPERATION)
220 @override
221 async def list_databases(
222 self,
223 limit: Optional[int] = None,
224 offset: Optional[int] = None,
225 tenant: str = DEFAULT_TENANT,
226 ) -> Sequence[Database]:
227 response = await self._make_request(
228 "get",
229 f"/tenants/{tenant}/databases",
230 params=BaseHTTPClient._clean_params(
231 {
232 "limit": limit,
233 "offset": offset,
234 }
235 ),
236 )
237
238 return [
239 Database(id=db["id"], name=db["name"], tenant=db["tenant"])
240 for db in response
241 ]
242
243 @trace_method("AsyncFastAPI.create_tenant", OpenTelemetryGranularity.OPERATION)
244 @override
245 async def create_tenant(self, name: str) -> None:
246 await self._make_request(
247 "post",
248 "/tenants",
249 json={"name": name},
250 )
251
252 @trace_method("AsyncFastAPI.get_tenant", OpenTelemetryGranularity.OPERATION)
253 @override
254 async def get_tenant(self, name: str) -> Tenant:
255 resp_json = await self._make_request(
256 "get",
257 "/tenants/" + name,
258 )
259
260 return Tenant(name=resp_json["name"])
261
262 @trace_method("AsyncFastAPI.get_user_identity", OpenTelemetryGranularity.OPERATION)
263 @override
264 async def get_user_identity(self) -> UserIdentity:
265 return UserIdentity(**(await self._make_request("get", "/auth/identity")))
266
267 @trace_method("AsyncFastAPI.list_collections", OpenTelemetryGranularity.OPERATION)
268 @override
269 async def list_collections(
270 self,
271 limit: Optional[int] = None,
272 offset: Optional[int] = None,
273 tenant: str = DEFAULT_TENANT,
274 database: str = DEFAULT_DATABASE,
275 ) -> Sequence[CollectionModel]:
276 resp_json = await self._make_request(
277 "get",
278 f"/tenants/{tenant}/databases/{database}/collections",
279 params=BaseHTTPClient._clean_params(
280 {
281 "limit": limit,
282 "offset": offset,
283 }
284 ),
285 )
286
287 models = [
288 CollectionModel.from_json(json_collection) for json_collection in resp_json
289 ]
290 return models
291
292 @trace_method("AsyncFastAPI.count_collections", OpenTelemetryGranularity.OPERATION)
293 @override
294 async def count_collections(
295 self, tenant: str = DEFAULT_TENANT, database: str = DEFAULT_DATABASE
296 ) -> int:
297 resp_json = await self._make_request(
298 "get",
299 f"/tenants/{tenant}/databases/{database}/collections_count",
300 )
301
302 return cast(int, resp_json)
303
304 @trace_method("AsyncFastAPI.create_collection", OpenTelemetryGranularity.OPERATION)
305 @override
306 async def create_collection(
307 self,
308 name: str,
309 schema: Optional[Schema] = None,
310 configuration: Optional[CreateCollectionConfiguration] = None,
311 metadata: Optional[CollectionMetadata] = None,
312 get_or_create: bool = False,
313 tenant: str = DEFAULT_TENANT,
314 database: str = DEFAULT_DATABASE,
315 ) -> CollectionModel:
316 """Creates a collection"""
317 config_json = (
318 create_collection_configuration_to_json(configuration, metadata)
319 if configuration
320 else None
321 )
322 serialized_schema = schema.serialize_to_json() if schema else None
323 resp_json = await self._make_request(
324 "post",
325 f"/tenants/{tenant}/databases/{database}/collections",
326 json={
327 "name": name,
328 "metadata": metadata,
329 "configuration": config_json,
330 "schema": serialized_schema,
331 "get_or_create": get_or_create,
332 },
333 )
334 model = CollectionModel.from_json(resp_json)
335
336 return model
337
338 @trace_method("AsyncFastAPI.get_collection", OpenTelemetryGranularity.OPERATION)
339 @override
340 async def get_collection(
341 self,
342 name: str,
343 tenant: str = DEFAULT_TENANT,
344 database: str = DEFAULT_DATABASE,
345 ) -> CollectionModel:
346 resp_json = await self._make_request(
347 "get",
348 f"/tenants/{tenant}/databases/{database}/collections/{name}",
349 )
350
351 model = CollectionModel.from_json(resp_json)
352
353 return model
354
355 @trace_method(
356 "AsyncFastAPI.get_collection_by_id", OpenTelemetryGranularity.OPERATION
357 )
358 @override
359 async def get_collection_by_id(
360 self,
361 collection_id: UUID,
362 tenant: str = DEFAULT_TENANT,
363 database: str = DEFAULT_DATABASE,
364 ) -> CollectionModel:
365 """Returns a collection by its ID"""
366 resp_json = await self._make_request(
367 "get",
368 f"/tenants/{tenant}/databases/{database}/collections/by-id/{collection_id}",
369 )
370
371 model = CollectionModel.from_json(resp_json)
372
373 return model
374
375 @trace_method(
376 "AsyncFastAPI.get_or_create_collection", OpenTelemetryGranularity.OPERATION
377 )
378 @override
379 async def get_or_create_collection(
380 self,
381 name: str,
382 schema: Optional[Schema] = None,
383 configuration: Optional[CreateCollectionConfiguration] = None,
384 metadata: Optional[CollectionMetadata] = None,
385 tenant: str = DEFAULT_TENANT,
386 database: str = DEFAULT_DATABASE,
387 ) -> CollectionModel:
388 return await self.create_collection(
389 name=name,
390 schema=schema,
391 configuration=configuration,
392 metadata=metadata,
393 get_or_create=True,
394 tenant=tenant,
395 database=database,
396 )
397
398 @trace_method("AsyncFastAPI._modify", OpenTelemetryGranularity.OPERATION)
399 @override
400 async def _modify(
401 self,
402 id: UUID,
403 new_name: Optional[str] = None,
404 new_metadata: Optional[CollectionMetadata] = None,
405 new_configuration: Optional[UpdateCollectionConfiguration] = None,
406 tenant: str = DEFAULT_TENANT,
407 database: str = DEFAULT_DATABASE,
408 ) -> None:
409 await self._make_request(
410 "put",
411 f"/tenants/{tenant}/databases/{database}/collections/{id}",
412 json={
413 "new_metadata": new_metadata,
414 "new_name": new_name,
415 "new_configuration": update_collection_configuration_to_json(
416 new_configuration
417 )
418 if new_configuration
419 else None,
420 },
421 )
422
423 @trace_method("AsyncFastAPI._fork", OpenTelemetryGranularity.OPERATION)
424 @override
425 async def _fork(
426 self,
427 collection_id: UUID,
428 new_name: str,
429 tenant: str = DEFAULT_TENANT,
430 database: str = DEFAULT_DATABASE,
431 ) -> CollectionModel:
432 resp_json = await self._make_request(
433 "post",
434 f"/tenants/{tenant}/databases/{database}/collections/{collection_id}/fork",
435 json={"new_name": new_name},
436 )
437 model = CollectionModel.from_json(resp_json)
438 return model
439
440 @trace_method("AsyncFastAPI._fork_count", OpenTelemetryGranularity.OPERATION)
441 @override
442 async def _fork_count(
443 self,
444 collection_id: UUID,
445 tenant: str = DEFAULT_TENANT,
446 database: str = DEFAULT_DATABASE,
447 ) -> int:
448 resp_json = await self._make_request(
449 "get",
450 f"/tenants/{tenant}/databases/{database}/collections/{collection_id}/fork_count",
451 )
452 return int(resp_json["count"])
453
454 @trace_method(
455 "AsyncFastAPI._get_indexing_status", OpenTelemetryGranularity.OPERATION
456 )
457 @override
458 async def _get_indexing_status(
459 self,
460 collection_id: UUID,
461 tenant: str = DEFAULT_TENANT,
462 database: str = DEFAULT_DATABASE,
463 ) -> IndexingStatus:
464 resp_json = await self._make_request(
465 "get",
466 f"/tenants/{tenant}/databases/{database}/collections/{collection_id}/indexing_status",
467 )
468 return IndexingStatus(
469 num_indexed_ops=resp_json["num_indexed_ops"],
470 num_unindexed_ops=resp_json["num_unindexed_ops"],
471 total_ops=resp_json["total_ops"],
472 op_indexing_progress=resp_json["op_indexing_progress"],
473 )
474
475 @trace_method("AsyncFastAPI._search", OpenTelemetryGranularity.OPERATION)
476 @override
477 async def _search(
478 self,
479 collection_id: UUID,
480 searches: List[Search],
481 tenant: str = DEFAULT_TENANT,
482 database: str = DEFAULT_DATABASE,
483 read_level: ReadLevel = ReadLevel.INDEX_AND_WAL,
484 ) -> SearchResult:
485 """Performs hybrid search on a collection"""
486 payload = {
487 "searches": [s.to_dict() for s in searches],
488 "read_level": read_level,
489 }
490
491 resp_json = await self._make_request(
492 "post",
493 f"/tenants/{tenant}/databases/{database}/collections/{collection_id}/search",
494 json=payload,
495 )
496
497 metadata_batches = resp_json.get("metadatas", None)
498 if metadata_batches is not None:
499 resp_json["metadatas"] = [
500 [
501 deserialize_metadata(metadata) if metadata is not None else None
502 for metadata in metadatas
503 ]
504 if metadatas is not None
505 else None
506 for metadatas in metadata_batches
507 ]
508
509 return SearchResult(resp_json)
510
511 @trace_method("AsyncFastAPI.delete_collection", OpenTelemetryGranularity.OPERATION)
512 @override
513 async def delete_collection(
514 self,
515 name: str,
516 tenant: str = DEFAULT_TENANT,
517 database: str = DEFAULT_DATABASE,
518 ) -> None:
519 await self._make_request(
520 "delete",
521 f"/tenants/{tenant}/databases/{database}/collections/{name}",
522 )
523
524 @trace_method("AsyncFastAPI._count", OpenTelemetryGranularity.OPERATION)
525 @override
526 async def _count(
527 self,
528 collection_id: UUID,
529 tenant: str = DEFAULT_TENANT,
530 database: str = DEFAULT_DATABASE,
531 read_level: ReadLevel = ReadLevel.INDEX_AND_WAL,
532 ) -> int:
533 """Returns the number of embeddings in the database"""
534 resp_json = await self._make_request(
535 "get",
536 f"/tenants/{tenant}/databases/{database}/collections/{collection_id}/count",
537 params={"read_level": read_level.value},
538 )
539
540 return cast(int, resp_json)
541
542 @trace_method("AsyncFastAPI._peek", OpenTelemetryGranularity.OPERATION)
543 @override
544 async def _peek(
545 self,
546 collection_id: UUID,
547 n: int = 10,
548 tenant: str = DEFAULT_TENANT,
549 database: str = DEFAULT_DATABASE,
550 ) -> GetResult:
551 resp = await self._get(
552 collection_id,
553 tenant=tenant,
554 database=database,
555 limit=n,
556 include=IncludeMetadataDocumentsEmbeddings,
557 )
558
559 return resp
560
561 @trace_method("AsyncFastAPI._get", OpenTelemetryGranularity.OPERATION)
562 @override
563 async def _get(
564 self,
565 collection_id: UUID,
566 ids: Optional[IDs] = None,
567 where: Optional[Where] = None,
568 limit: Optional[int] = None,
569 offset: Optional[int] = None,
570 where_document: Optional[WhereDocument] = None,
571 include: Include = IncludeMetadataDocuments,
572 tenant: str = DEFAULT_TENANT,
573 database: str = DEFAULT_DATABASE,
574 ) -> GetResult:
575 # Servers do not support the "data" include, as that is hydrated on the client side
576 filtered_include = [i for i in include if i != "data"]
577
578 resp_json = await self._make_request(
579 "post",
580 f"/tenants/{tenant}/databases/{database}/collections/{collection_id}/get",
581 json={
582 "ids": ids,
583 "where": where,
584 "limit": limit,
585 "offset": offset,
586 "where_document": where_document,
587 "include": filtered_include,
588 },
589 )
590
591 metadatas = resp_json.get("metadatas", None)
592 if metadatas is not None:
593 metadatas = [
594 deserialize_metadata(metadata) if metadata is not None else None
595 for metadata in metadatas
596 ]
597
598 return GetResult(
599 ids=resp_json["ids"],
600 embeddings=resp_json.get("embeddings", None),
601 metadatas=metadatas,
602 documents=resp_json.get("documents", None),
603 data=None,
604 uris=resp_json.get("uris", None),
605 included=include,
606 )
607
608 @trace_method("AsyncFastAPI._delete", OpenTelemetryGranularity.OPERATION)
609 @override
610 async def _delete(
611 self,
612 collection_id: UUID,
613 ids: Optional[IDs] = None,
614 where: Optional[Where] = None,
615 where_document: Optional[WhereDocument] = None,
616 limit: Optional[int] = None,
617 tenant: str = DEFAULT_TENANT,
618 database: str = DEFAULT_DATABASE,
619 ) -> DeleteResult:
620 body: dict = {"where": where, "ids": ids, "where_document": where_document}
621 if limit is not None:
622 body["limit"] = limit
623 resp = await self._make_request(
624 "post",
625 f"/tenants/{tenant}/databases/{database}/collections/{collection_id}/delete",
626 json=body,
627 )
628 return DeleteResult(deleted=resp.get("deleted", 0) if resp else 0)
629
630 @trace_method("AsyncFastAPI._submit_batch", OpenTelemetryGranularity.ALL)
631 async def _submit_batch(
632 self,
633 batch: Tuple[
634 IDs,
635 Optional[Embeddings],
636 Optional[Metadatas],
637 Optional[Documents],
638 Optional[URIs],
639 ],
640 url: str,
641 ) -> Any:
642 """
643 Submits a batch of embeddings to the database
644 """
645 supports_base64_encoding = await self.supports_base64_encoding()
646
647 serialized_metadatas = None
648 if batch[2] is not None:
649 serialized_metadatas = [
650 serialize_metadata(metadata) if metadata is not None else None
651 for metadata in batch[2]
652 ]
653
654 data = {
655 "ids": batch[0],
656 "embeddings": optional_embeddings_to_base64_strings(batch[1])
657 if supports_base64_encoding
658 else batch[1],
659 "metadatas": serialized_metadatas,
660 "documents": batch[3],
661 "uris": batch[4],
662 }
663
664 return await self._make_request(
665 "post",
666 url,
667 json=data,
668 )
669
670 @trace_method("AsyncFastAPI._add", OpenTelemetryGranularity.ALL)
671 @override
672 async def _add(
673 self,
674 ids: IDs,
675 collection_id: UUID,
676 embeddings: Embeddings,
677 metadatas: Optional[Metadatas] = None,
678 documents: Optional[Documents] = None,
679 uris: Optional[URIs] = None,
680 tenant: str = DEFAULT_TENANT,
681 database: str = DEFAULT_DATABASE,
682 ) -> bool:
683 batch = (
684 ids,
685 embeddings,
686 metadatas,
687 documents,
688 uris,
689 )
690 validate_batch(batch, {"max_batch_size": await self.get_max_batch_size()})
691 await self._submit_batch(
692 batch,
693 f"/tenants/{tenant}/databases/{database}/collections/{str(collection_id)}/add",
694 )
695 return True
696
697 @trace_method("AsyncFastAPI._update", OpenTelemetryGranularity.ALL)
698 @override
699 async def _update(
700 self,
701 collection_id: UUID,
702 ids: IDs,
703 embeddings: Optional[Embeddings] = None,
704 metadatas: Optional[Metadatas] = None,
705 documents: Optional[Documents] = None,
706 uris: Optional[URIs] = None,
707 tenant: str = DEFAULT_TENANT,
708 database: str = DEFAULT_DATABASE,
709 ) -> bool:
710 batch = (
711 ids,
712 embeddings if embeddings is not None else None,
713 metadatas,
714 documents,
715 uris,
716 )
717 validate_batch(batch, {"max_batch_size": await self.get_max_batch_size()})
718
719 await self._submit_batch(
720 batch,
721 f"/tenants/{tenant}/databases/{database}/collections/{str(collection_id)}/update",
722 )
723
724 return True
725
726 @trace_method("AsyncFastAPI._upsert", OpenTelemetryGranularity.ALL)
727 @override
728 async def _upsert(
729 self,
730 collection_id: UUID,
731 ids: IDs,
732 embeddings: Embeddings,
733 metadatas: Optional[Metadatas] = None,
734 documents: Optional[Documents] = None,
735 uris: Optional[URIs] = None,
736 tenant: str = DEFAULT_TENANT,
737 database: str = DEFAULT_DATABASE,
738 ) -> bool:
739 batch = (
740 ids,
741 embeddings,
742 metadatas,
743 documents,
744 uris,
745 )
746 validate_batch(batch, {"max_batch_size": await self.get_max_batch_size()})
747 await self._submit_batch(
748 batch,
749 f"/tenants/{tenant}/databases/{database}/collections/{str(collection_id)}/upsert",
750 )
751 return True
752
753 @trace_method("AsyncFastAPI._query", OpenTelemetryGranularity.ALL)
754 @override
755 async def _query(
756 self,
757 collection_id: UUID,
758 query_embeddings: Embeddings,
759 ids: Optional[IDs] = None,
760 n_results: int = 10,
761 where: Optional[Where] = None,
762 where_document: Optional[WhereDocument] = None,
763 include: Include = IncludeMetadataDocumentsDistances,
764 tenant: str = DEFAULT_TENANT,
765 database: str = DEFAULT_DATABASE,
766 ) -> QueryResult:
767 # Servers do not support the "data" include, as that is hydrated on the client side
768 filtered_include = [i for i in include if i != "data"]
769
770 resp_json = await self._make_request(
771 "post",
772 f"/tenants/{tenant}/databases/{database}/collections/{collection_id}/query",
773 json={
774 "ids": ids,
775 "query_embeddings": convert_np_embeddings_to_list(query_embeddings)
776 if query_embeddings is not None
777 else None,
778 "n_results": n_results,
779 "where": where,
780 "where_document": where_document,
781 "include": filtered_include,
782 },
783 )
784
785 metadata_batches = resp_json.get("metadatas", None)
786 if metadata_batches is not None:
787 metadata_batches = [
788 [
789 deserialize_metadata(metadata) if metadata is not None else None
790 for metadata in metadatas
791 ]
792 if metadatas is not None
793 else None
794 for metadatas in metadata_batches
795 ]
796
797 return QueryResult(
798 ids=resp_json["ids"],
799 distances=resp_json.get("distances", None),
800 embeddings=resp_json.get("embeddings", None),
801 metadatas=metadata_batches,
802 documents=resp_json.get("documents", None),
803 uris=resp_json.get("uris", None),
804 data=None,
805 included=include,
806 )
807
808 @trace_method("AsyncFastAPI.reset", OpenTelemetryGranularity.ALL)
809 @override
810 async def reset(self) -> bool:
811 resp_json = await self._make_request("post", "/reset")
812 return cast(bool, resp_json)
813
814 @trace_method("AsyncFastAPI.get_version", OpenTelemetryGranularity.OPERATION)
815 @override
816 async def get_version(self) -> str:
817 resp_json = await self._make_request("get", "/version")
818 return cast(str, resp_json)
819
820 @override
821 def get_settings(self) -> Settings:
822 return self._settings
823
824 @trace_method(
825 "AsyncFastAPI.get_pre_flight_checks", OpenTelemetryGranularity.OPERATION
826 )
827 async def get_pre_flight_checks(self) -> Any:
828 if self.pre_flight_checks is None:
829 resp_json = await self._make_request("get", "/pre-flight-checks")
830 self.pre_flight_checks = resp_json
831 return self.pre_flight_checks
832
833 @trace_method(
834 "AsyncFastAPI.supports_base64_encoding", OpenTelemetryGranularity.OPERATION
835 )
836 async def supports_base64_encoding(self) -> bool:
837 pre_flight_checks = await self.get_pre_flight_checks()
838 b64_encoding_enabled = cast(
839 bool, pre_flight_checks.get("supports_base64_encoding", False)
840 )
841 return b64_encoding_enabled
842
843 @trace_method("AsyncFastAPI.get_max_batch_size", OpenTelemetryGranularity.OPERATION)
844 @override
845 async def get_max_batch_size(self) -> int:
846 pre_flight_checks = await self.get_pre_flight_checks()
847 max_batch_size = cast(int, pre_flight_checks.get("max_batch_size", -1))
848 return max_batch_size
849 