Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
distributed.py243 linesDownload Raw Back to executor
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 
codekingpro/portable-devtools · Team Ai