Team Ai
Datasetpublic

codekingpro/portable-devtools

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