Team Ai
Datasetpublic

codekingpro/portable-devtools

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