Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
client.py559 linesDownload Raw Back to grpc
1from typing import List, Optional, Sequence, Tuple, Union, cast
2from uuid import UUID
3from overrides import overrides
4from chromadb.api.collection_configuration import (
5    CreateCollectionConfiguration,
6    create_collection_configuration_to_json_str,
7    UpdateCollectionConfiguration,
8    update_collection_configuration_to_json_str,
9    CollectionMetadata,
10)
11from chromadb.api.types import Schema
12from chromadb.config import DEFAULT_DATABASE, DEFAULT_TENANT, System, logger
13from chromadb.db.system import SysDB
14from chromadb.errors import NotFoundError, UniqueConstraintError, InternalError
15from chromadb.proto.convert import (
16    from_proto_collection,
17    from_proto_segment,
18    to_proto_update_metadata,
19    to_proto_segment,
20    to_proto_segment_scope,
21)
22from chromadb.proto.coordinator_pb2 import (
23    CreateCollectionRequest,
24    CreateDatabaseRequest,
25    CreateSegmentRequest,
26    CreateTenantRequest,
27    CountCollectionsRequest,
28    CountCollectionsResponse,
29    DeleteCollectionRequest,
30    DeleteDatabaseRequest,
31    DeleteSegmentRequest,
32    GetCollectionsRequest,
33    GetCollectionsResponse,
34    GetCollectionSizeRequest,
35    GetCollectionSizeResponse,
36    GetCollectionWithSegmentsRequest,
37    GetCollectionWithSegmentsResponse,
38    GetDatabaseRequest,
39    GetSegmentsRequest,
40    GetTenantRequest,
41    ListDatabasesRequest,
42    UpdateCollectionRequest,
43    UpdateSegmentRequest,
44)
45from chromadb.proto.coordinator_pb2_grpc import SysDBStub
46from chromadb.proto.utils import RetryOnRpcErrorClientInterceptor
47from chromadb.telemetry.opentelemetry.grpc import OtelInterceptor
48from chromadb.telemetry.opentelemetry import (
49    OpenTelemetryGranularity,
50    trace_method,
51)
52from chromadb.types import (
53    Collection,
54    CollectionAndSegments,
55    Database,
56    Metadata,
57    OptionalArgument,
58    Segment,
59    SegmentScope,
60    Tenant,
61    Unspecified,
62    UpdateMetadata,
63)
64from google.protobuf.empty_pb2 import Empty
65import grpc
66
67
68class GrpcSysDB(SysDB):
69    """A gRPC implementation of the SysDB. In the distributed system, the SysDB is also
70    called the 'Coordinator'. This implementation is used by Chroma frontend servers
71    to call a remote SysDB (Coordinator) service."""
72
73    _sys_db_stub: SysDBStub
74    _channel: grpc.Channel
75    _coordinator_url: str
76    _coordinator_port: int
77    _request_timeout_seconds: int
78
79    def __init__(self, system: System):
80        self._coordinator_url = system.settings.require("chroma_coordinator_host")
81        # TODO: break out coordinator_port into a separate setting?
82        self._coordinator_port = system.settings.require("chroma_server_grpc_port")
83        self._request_timeout_seconds = system.settings.require(
84            "chroma_sysdb_request_timeout_seconds"
85        )
86        return super().__init__(system)
87
88    @overrides
89    def start(self) -> None:
90        self._channel = grpc.insecure_channel(
91            f"{self._coordinator_url}:{self._coordinator_port}",
92            options=[("grpc.max_concurrent_streams", 1000)],
93        )
94        interceptors = [OtelInterceptor(), RetryOnRpcErrorClientInterceptor()]
95        self._channel = grpc.intercept_channel(self._channel, *interceptors)
96        self._sys_db_stub = SysDBStub(self._channel)  # type: ignore
97        return super().start()
98
99    @overrides
100    def stop(self) -> None:
101        self._channel.close()
102        return super().stop()
103
104    @overrides
105    def reset_state(self) -> None:
106        self._sys_db_stub.ResetState(Empty())
107        return super().reset_state()
108
109    @overrides
110    def create_database(
111        self, id: UUID, name: str, tenant: str = DEFAULT_TENANT
112    ) -> None:
113        try:
114            request = CreateDatabaseRequest(id=id.hex, name=name, tenant=tenant)
115            response = self._sys_db_stub.CreateDatabase(
116                request, timeout=self._request_timeout_seconds
117            )
118        except grpc.RpcError as e:
119            logger.info(
120                f"Failed to create database name {name} and database id {id} for tenant {tenant} due to error: {e}"
121            )
122            if e.code() == grpc.StatusCode.ALREADY_EXISTS:
123                raise UniqueConstraintError()
124            raise InternalError()
125
126    @overrides
127    def get_database(self, name: str, tenant: str = DEFAULT_TENANT) -> Database:
128        try:
129            request = GetDatabaseRequest(name=name, tenant=tenant)
130            response = self._sys_db_stub.GetDatabase(
131                request, timeout=self._request_timeout_seconds
132            )
133            return Database(
134                id=UUID(hex=response.database.id),
135                name=response.database.name,
136                tenant=response.database.tenant,
137            )
138        except grpc.RpcError as e:
139            logger.info(
140                f"Failed to get database {name} for tenant {tenant} due to error: {e}"
141            )
142            if e.code() == grpc.StatusCode.NOT_FOUND:
143                raise NotFoundError()
144            raise InternalError()
145
146    @overrides
147    def delete_database(self, name: str, tenant: str = DEFAULT_TENANT) -> None:
148        try:
149            request = DeleteDatabaseRequest(name=name, tenant=tenant)
150            self._sys_db_stub.DeleteDatabase(
151                request, timeout=self._request_timeout_seconds
152            )
153        except grpc.RpcError as e:
154            logger.info(
155                f"Failed to delete database {name} for tenant {tenant} due to error: {e}"
156            )
157            if e.code() == grpc.StatusCode.NOT_FOUND:
158                raise NotFoundError()
159            raise InternalError
160
161    @overrides
162    def list_databases(
163        self,
164        limit: Optional[int] = None,
165        offset: Optional[int] = None,
166        tenant: str = DEFAULT_TENANT,
167    ) -> Sequence[Database]:
168        try:
169            request = ListDatabasesRequest(limit=limit, offset=offset, tenant=tenant)
170            response = self._sys_db_stub.ListDatabases(
171                request, timeout=self._request_timeout_seconds
172            )
173            results: List[Database] = []
174            for proto_database in response.databases:
175                results.append(
176                    Database(
177                        id=UUID(hex=proto_database.id),
178                        name=proto_database.name,
179                        tenant=proto_database.tenant,
180                    )
181                )
182            return results
183        except grpc.RpcError as e:
184            logger.info(
185                f"Failed to list databases for tenant {tenant} due to error: {e}"
186            )
187            raise InternalError()
188
189    @overrides
190    def create_tenant(self, name: str) -> None:
191        try:
192            request = CreateTenantRequest(name=name)
193            response = self._sys_db_stub.CreateTenant(
194                request, timeout=self._request_timeout_seconds
195            )
196        except grpc.RpcError as e:
197            logger.info(f"Failed to create tenant {name} due to error: {e}")
198            if e.code() == grpc.StatusCode.ALREADY_EXISTS:
199                raise UniqueConstraintError()
200            raise InternalError()
201
202    @overrides
203    def get_tenant(self, name: str) -> Tenant:
204        try:
205            request = GetTenantRequest(name=name)
206            response = self._sys_db_stub.GetTenant(
207                request, timeout=self._request_timeout_seconds
208            )
209            return Tenant(
210                name=response.tenant.name,
211            )
212        except grpc.RpcError as e:
213            logger.info(f"Failed to get tenant {name} due to error: {e}")
214            if e.code() == grpc.StatusCode.NOT_FOUND:
215                raise NotFoundError()
216            raise InternalError()
217
218    @overrides
219    def create_segment(self, segment: Segment) -> None:
220        try:
221            proto_segment = to_proto_segment(segment)
222            request = CreateSegmentRequest(
223                segment=proto_segment,
224            )
225            response = self._sys_db_stub.CreateSegment(
226                request, timeout=self._request_timeout_seconds
227            )
228        except grpc.RpcError as e:
229            logger.info(f"Failed to create segment {segment}, error: {e}")
230            if e.code() == grpc.StatusCode.ALREADY_EXISTS:
231                raise UniqueConstraintError()
232            raise InternalError()
233
234    @overrides
235    def delete_segment(self, collection: UUID, id: UUID) -> None:
236        try:
237            request = DeleteSegmentRequest(
238                id=id.hex,
239                collection=collection.hex,
240            )
241            response = self._sys_db_stub.DeleteSegment(
242                request, timeout=self._request_timeout_seconds
243            )
244        except grpc.RpcError as e:
245            logger.info(
246                f"Failed to delete segment with id {id} for collection {collection} due to error: {e}"
247            )
248            if e.code() == grpc.StatusCode.NOT_FOUND:
249                raise NotFoundError()
250            raise InternalError()
251
252    @overrides
253    def get_segments(
254        self,
255        collection: UUID,
256        id: Optional[UUID] = None,
257        type: Optional[str] = None,
258        scope: Optional[SegmentScope] = None,
259    ) -> Sequence[Segment]:
260        try:
261            request = GetSegmentsRequest(
262                id=id.hex if id else None,
263                type=type,
264                scope=to_proto_segment_scope(scope) if scope else None,
265                collection=collection.hex,
266            )
267            response = self._sys_db_stub.GetSegments(
268                request, timeout=self._request_timeout_seconds
269            )
270            results: List[Segment] = []
271            for proto_segment in response.segments:
272                segment = from_proto_segment(proto_segment)
273                results.append(segment)
274            return results
275        except grpc.RpcError as e:
276            logger.info(
277                f"Failed to get segment id {id}, type {type}, scope {scope} for collection {collection} due to error: {e}"
278            )
279            raise InternalError()
280
281    @overrides
282    def update_segment(
283        self,
284        collection: UUID,
285        id: UUID,
286        metadata: OptionalArgument[Optional[UpdateMetadata]] = Unspecified(),
287    ) -> None:
288        try:
289            write_metadata = None
290            if metadata != Unspecified():
291                write_metadata = cast(Union[UpdateMetadata, None], metadata)
292
293            request = UpdateSegmentRequest(
294                id=id.hex,
295                collection=collection.hex,
296                metadata=to_proto_update_metadata(write_metadata)
297                if write_metadata
298                else None,
299            )
300
301            if metadata is None:
302                request.ClearField("metadata")
303                request.reset_metadata = True
304
305            self._sys_db_stub.UpdateSegment(
306                request, timeout=self._request_timeout_seconds
307            )
308        except grpc.RpcError as e:
309            logger.info(
310                f"Failed to update segment with id {id} for collection {collection}, error: {e}"
311            )
312            raise InternalError()
313
314    @overrides
315    def create_collection(
316        self,
317        id: UUID,
318        name: str,
319        schema: Optional[Schema],
320        configuration: CreateCollectionConfiguration,
321        segments: Sequence[Segment],
322        metadata: Optional[Metadata] = None,
323        dimension: Optional[int] = None,
324        get_or_create: bool = False,
325        tenant: str = DEFAULT_TENANT,
326        database: str = DEFAULT_DATABASE,
327    ) -> Tuple[Collection, bool]:
328        try:
329            request = CreateCollectionRequest(
330                id=id.hex,
331                name=name,
332                configuration_json_str=create_collection_configuration_to_json_str(
333                    configuration, cast(CollectionMetadata, metadata)
334                ),
335                metadata=to_proto_update_metadata(metadata) if metadata else None,
336                dimension=dimension,
337                get_or_create=get_or_create,
338                tenant=tenant,
339                database=database,
340                segments=[to_proto_segment(segment) for segment in segments],
341            )
342            response = self._sys_db_stub.CreateCollection(
343                request, timeout=self._request_timeout_seconds
344            )
345            collection = from_proto_collection(response.collection)
346            return collection, response.created
347        except grpc.RpcError as e:
348            logger.error(
349                f"Failed to create collection id {id}, name {name} for database {database} and tenant {tenant} due to error: {e}"
350            )
351            if e.code() == grpc.StatusCode.ALREADY_EXISTS:
352                raise UniqueConstraintError()
353            raise InternalError()
354
355    @overrides
356    def delete_collection(
357        self,
358        id: UUID,
359        tenant: str = DEFAULT_TENANT,
360        database: str = DEFAULT_DATABASE,
361    ) -> None:
362        try:
363            request = DeleteCollectionRequest(
364                id=id.hex,
365                tenant=tenant,
366                database=database,
367            )
368            response = self._sys_db_stub.DeleteCollection(
369                request, timeout=self._request_timeout_seconds
370            )
371        except grpc.RpcError as e:
372            logger.error(
373                f"Failed to delete collection id {id} for database {database} and tenant {tenant} due to error: {e}"
374            )
375            e = cast(grpc.Call, e)
376            logger.error(
377                f"Error code: {e.code()}, NotFoundError: {grpc.StatusCode.NOT_FOUND}"
378            )
379            if e.code() == grpc.StatusCode.NOT_FOUND:
380                raise NotFoundError()
381            raise InternalError()
382
383    @overrides
384    def get_collections(
385        self,
386        id: Optional[UUID] = None,
387        name: Optional[str] = None,
388        tenant: str = DEFAULT_TENANT,
389        database: str = DEFAULT_DATABASE,
390        limit: Optional[int] = None,
391        offset: Optional[int] = None,
392    ) -> Sequence[Collection]:
393        try:
394            # TODO: implement limit and offset in the gRPC service
395            request = None
396            if id is not None:
397                request = GetCollectionsRequest(
398                    id=id.hex,
399                    limit=limit,
400                    offset=offset,
401                )
402            if name is not None:
403                if tenant is None and database is None:
404                    raise ValueError(
405                        "If name is specified, tenant and database must also be specified in order to uniquely identify the collection"
406                    )
407                request = GetCollectionsRequest(
408                    name=name,
409                    tenant=tenant,
410                    database=database,
411                    limit=limit,
412                    offset=offset,
413                )
414            if id is None and name is None:
415                request = GetCollectionsRequest(
416                    tenant=tenant,
417                    database=database,
418                    limit=limit,
419                    offset=offset,
420                )
421            response: GetCollectionsResponse = self._sys_db_stub.GetCollections(
422                request, timeout=self._request_timeout_seconds
423            )
424            results: List[Collection] = []
425            for collection in response.collections:
426                results.append(from_proto_collection(collection))
427            return results
428        except grpc.RpcError as e:
429            logger.error(
430                f"Failed to get collections with id {id}, name {name}, tenant {tenant}, database {database} due to error: {e}"
431            )
432            raise InternalError()
433
434    @overrides
435    def count_collections(
436        self,
437        tenant: str = DEFAULT_TENANT,
438        database: Optional[str] = None,
439    ) -> int:
440        try:
441            if database is None or database == "":
442                request = CountCollectionsRequest(tenant=tenant)
443                response: CountCollectionsResponse = self._sys_db_stub.CountCollections(
444                    request
445                )
446                return response.count
447            else:
448                request = CountCollectionsRequest(
449                    tenant=tenant,
450                    database=database,
451                )
452                response: CountCollectionsResponse = self._sys_db_stub.CountCollections(
453                    request
454                )
455                return response.count
456        except grpc.RpcError as e:
457            logger.error(f"Failed to count collections due to error: {e}")
458            raise InternalError()
459
460    @overrides
461    def get_collection_size(self, id: UUID) -> int:
462        try:
463            request = GetCollectionSizeRequest(id=id.hex)
464            response: GetCollectionSizeResponse = self._sys_db_stub.GetCollectionSize(
465                request
466            )
467            return response.total_records_post_compaction
468        except grpc.RpcError as e:
469            logger.error(f"Failed to get collection {id} size due to error: {e}")
470            raise InternalError()
471
472    @trace_method(
473        "SysDB.get_collection_with_segments", OpenTelemetryGranularity.OPERATION
474    )
475    @overrides
476    def get_collection_with_segments(
477        self, collection_id: UUID
478    ) -> CollectionAndSegments:
479        try:
480            request = GetCollectionWithSegmentsRequest(id=collection_id.hex)
481            response: GetCollectionWithSegmentsResponse = (
482                self._sys_db_stub.GetCollectionWithSegments(request)
483            )
484            return CollectionAndSegments(
485                collection=from_proto_collection(response.collection),
486                segments=[from_proto_segment(segment) for segment in response.segments],
487            )
488        except grpc.RpcError as e:
489            if e.code() == grpc.StatusCode.NOT_FOUND:
490                raise NotFoundError()
491            logger.error(
492                f"Failed to get collection {collection_id} and its segments due to error: {e}"
493            )
494            raise InternalError()
495
496    @overrides
497    def update_collection(
498        self,
499        id: UUID,
500        name: OptionalArgument[str] = Unspecified(),
501        dimension: OptionalArgument[Optional[int]] = Unspecified(),
502        metadata: OptionalArgument[Optional[UpdateMetadata]] = Unspecified(),
503        configuration: OptionalArgument[
504            Optional[UpdateCollectionConfiguration]
505        ] = Unspecified(),
506    ) -> None:
507        try:
508            write_name = None
509            if name != Unspecified():
510                write_name = cast(str, name)
511
512            write_dimension = None
513            if dimension != Unspecified():
514                write_dimension = cast(Union[int, None], dimension)
515
516            write_metadata = None
517            if metadata != Unspecified():
518                write_metadata = cast(Union[UpdateMetadata, None], metadata)
519
520            write_configuration = None
521            if configuration != Unspecified():
522                write_configuration = cast(
523                    Union[UpdateCollectionConfiguration, None], configuration
524                )
525
526            request = UpdateCollectionRequest(
527                id=id.hex,
528                name=write_name,
529                dimension=write_dimension,
530                metadata=to_proto_update_metadata(write_metadata)
531                if write_metadata
532                else None,
533                configuration_json_str=update_collection_configuration_to_json_str(
534                    write_configuration
535                )
536                if write_configuration
537                else None,
538            )
539            if metadata is None:
540                request.ClearField("metadata")
541                request.reset_metadata = True
542
543            response = self._sys_db_stub.UpdateCollection(
544                request, timeout=self._request_timeout_seconds
545            )
546        except grpc.RpcError as e:
547            e = cast(grpc.Call, e)
548            logger.error(
549                f"Failed to update collection id {id}, name {name} due to error: {e}"
550            )
551            if e.code() == grpc.StatusCode.NOT_FOUND:
552                raise NotFoundError()
553            if e.code() == grpc.StatusCode.ALREADY_EXISTS:
554                raise UniqueConstraintError()
555            raise InternalError()
556
557    def reset_and_wait_for_ready(self) -> None:
558        self._sys_db_stub.ResetState(Empty(), wait_for_ready=True)
559 
codekingpro/portable-devtools · Team Ai