codekingpro/portable-devtools
115k
1import threading
2import random
3from typing import Callable, Dict, List, Optional, TypeVar
4import grpc
5from overrides import overrides
6from chromadb.api.types import GetResult, Metadata, QueryResult
7from chromadb.config import System
8from chromadb.execution.executor.abstract import Executor
9from chromadb.execution.expression.operator import Scan
10from chromadb.execution.expression.plan import CountPlan, GetPlan, KNNPlan
11from chromadb.proto import convert
12from chromadb.proto.query_executor_pb2_grpc import QueryExecutorStub
13from chromadb.segment.impl.manager.distributed import DistributedSegmentManager
14from chromadb.telemetry.opentelemetry.grpc import OtelInterceptor
15from tenacity import (
16 RetryCallState,
17 Retrying,
18 stop_after_attempt,
19 wait_exponential_jitter,
20 retry_if_exception,
21)
22from opentelemetry.trace import Span
23
24
25def _clean_metadata(metadata: Optional[Metadata]) -> Optional[Metadata]:
26 """Remove any chroma-specific metadata keys that the client shouldn't see from a metadata map."""
27 if not metadata:
28 return None
29 result = {}
30 for k, v in metadata.items():
31 if not k.startswith("chroma:"):
32 result[k] = v
33 if len(result) == 0:
34 return None
35 return result
36
37
38def _uri(metadata: Optional[Metadata]) -> Optional[str]:
39 """Retrieve the uri (if any) from a Metadata map"""
40
41 if metadata and "chroma:uri" in metadata:
42 return str(metadata["chroma:uri"])
43 return None
44
45
46# Type variables for input and output types of the round-robin retry function
47I = TypeVar("I") # noqa: E741
48O = TypeVar("O") # noqa: E741
49
50
51class DistributedExecutor(Executor):
52 _mtx: threading.Lock
53 _grpc_stub_pool: Dict[str, QueryExecutorStub]
54 _manager: DistributedSegmentManager
55 _request_timeout_seconds: int
56 _query_replication_factor: int
57
58 def __init__(self, system: System):
59 super().__init__(system)
60 self._mtx = threading.Lock()
61 self._grpc_stub_pool = {}
62 self._manager = self.require(DistributedSegmentManager)
63 self._request_timeout_seconds = system.settings.require(
64 "chroma_query_request_timeout_seconds"
65 )
66 self._query_replication_factor = system.settings.require(
67 "chroma_query_replication_factor"
68 )
69
70 def _round_robin_retry(self, funcs: List[Callable[[I], O]], args: I) -> O:
71 """
72 Retry a list of functions in a round-robin fashion until one of them succeeds.
73
74 funcs: List of functions to retry
75 args: Arguments to pass to each function
76
77 """
78 attempt_count = 0
79 sleep_span: Optional[Span] = None
80
81 def before_sleep(_: RetryCallState) -> None:
82 # HACK(hammadb) 1/14/2024 - this is a hack to avoid the fact that tracer is not yet available and there are boot order issues
83 # This should really use our component system to get the tracer. Since our grpc utils use this pattern
84 # we are copying it here. This should be removed once we have a better way to get the tracer
85 from chromadb.telemetry.opentelemetry import tracer
86
87 nonlocal sleep_span
88 if tracer is not None:
89 sleep_span = tracer.start_span("Waiting to retry RPC")
90
91 for attempt in Retrying(
92 stop=stop_after_attempt(5),
93 wait=wait_exponential_jitter(0.1, jitter=0.1),
94 reraise=True,
95 retry=retry_if_exception(
96 lambda x: isinstance(x, grpc.RpcError)
97 and x.code() in [grpc.StatusCode.UNAVAILABLE, grpc.StatusCode.UNKNOWN]
98 ),
99 before_sleep=before_sleep,
100 ):
101 if sleep_span is not None:
102 sleep_span.end()
103 sleep_span = None
104
105 with attempt:
106 return funcs[attempt_count % len(funcs)](args)
107 attempt_count += 1
108
109 # NOTE(hammadb) because Retrying() will always either return or raise an exception, this line should never be reached
110 raise Exception("Unreachable code error - should never reach here")
111
112 @overrides
113 def count(self, plan: CountPlan) -> int:
114 endpoints = self._get_grpc_endpoints(plan.scan)
115 count_funcs = [self._get_stub(endpoint).Count for endpoint in endpoints]
116 count_result = self._round_robin_retry(
117 count_funcs, convert.to_proto_count_plan(plan)
118 )
119 return convert.from_proto_count_result(count_result)
120
121 @overrides
122 def get(self, plan: GetPlan) -> GetResult:
123 endpoints = self._get_grpc_endpoints(plan.scan)
124 get_funcs = [self._get_stub(endpoint).Get for endpoint in endpoints]
125 get_result = self._round_robin_retry(get_funcs, convert.to_proto_get_plan(plan))
126 records = convert.from_proto_get_result(get_result)
127
128 ids = [record["id"] for record in records]
129 embeddings = (
130 [record["embedding"] for record in records]
131 if plan.projection.embedding
132 else None
133 )
134 documents = (
135 [record["document"] for record in records]
136 if plan.projection.document
137 else None
138 )
139 uris = (
140 [_uri(record["metadata"]) for record in records]
141 if plan.projection.uri
142 else None
143 )
144 metadatas = (
145 [_clean_metadata(record["metadata"]) for record in records]
146 if plan.projection.metadata
147 else None
148 )
149
150 # TODO: Fix typing
151 return GetResult(
152 ids=ids,
153 embeddings=embeddings, # type: ignore[typeddict-item]
154 documents=documents, # type: ignore[typeddict-item]
155 uris=uris, # type: ignore[typeddict-item]
156 data=None,
157 metadatas=metadatas, # type: ignore[typeddict-item]
158 included=plan.projection.included,
159 )
160
161 @overrides
162 def knn(self, plan: KNNPlan) -> QueryResult:
163 endpoints = self._get_grpc_endpoints(plan.scan)
164 knn_funcs = [self._get_stub(endpoint).KNN for endpoint in endpoints]
165 knn_result = self._round_robin_retry(knn_funcs, convert.to_proto_knn_plan(plan))
166 results = convert.from_proto_knn_batch_result(knn_result)
167
168 ids = [[record["record"]["id"] for record in records] for records in results]
169 embeddings = (
170 [
171 [record["record"]["embedding"] for record in records]
172 for records in results
173 ]
174 if plan.projection.embedding
175 else None
176 )
177 documents = (
178 [
179 [record["record"]["document"] for record in records]
180 for records in results
181 ]
182 if plan.projection.document
183 else None
184 )
185 uris = (
186 [
187 [_uri(record["record"]["metadata"]) for record in records]
188 for records in results
189 ]
190 if plan.projection.uri
191 else None
192 )
193 metadatas = (
194 [
195 [_clean_metadata(record["record"]["metadata"]) for record in records]
196 for records in results
197 ]
198 if plan.projection.metadata
199 else None
200 )
201 distances = (
202 [[record["distance"] for record in records] for records in results]
203 if plan.projection.rank
204 else None
205 )
206
207 # TODO: Fix typing
208 return QueryResult(
209 ids=ids,
210 embeddings=embeddings, # type: ignore[typeddict-item]
211 documents=documents, # type: ignore[typeddict-item]
212 uris=uris, # type: ignore[typeddict-item]
213 data=None,
214 metadatas=metadatas, # type: ignore[typeddict-item]
215 distances=distances, # type: ignore[typeddict-item]
216 included=plan.projection.included,
217 )
218
219 def _get_grpc_endpoints(self, scan: Scan) -> List[str]:
220 # Since grpc endpoint is endpoint is determined by collection uuid,
221 # the endpoint should be the same for all segments of the same collection
222 grpc_urls = self._manager.get_endpoints(
223 scan.record, self._query_replication_factor
224 )
225 # Shuffle the grpc urls to distribute the load evenly
226 random.shuffle(grpc_urls)
227 return grpc_urls
228
229 def _get_stub(self, grpc_url: str) -> QueryExecutorStub:
230 with self._mtx:
231 if grpc_url not in self._grpc_stub_pool:
232 channel = grpc.insecure_channel(
233 grpc_url,
234 options=[
235 ("grpc.max_concurrent_streams", 1000),
236 ("grpc.max_receive_message_length", 32000000), # 32 MB
237 ],
238 )
239 interceptors = [OtelInterceptor()]
240 channel = grpc.intercept_channel(channel, *interceptors)
241 self._grpc_stub_pool[grpc_url] = QueryExecutorStub(channel)
242 return self._grpc_stub_pool[grpc_url]
243 