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