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