Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
segment_directory.py338 linesDownload Raw Back to distributed
1import threading
2import time
3from typing import Any, Callable, Dict, List, Optional, cast
4from kubernetes import client, config, watch
5from kubernetes.client.rest import ApiException
6from overrides import EnforceOverrides, override
7from chromadb.config import RoutingMode, System
8from chromadb.segment.distributed import (
9    Member,
10    Memberlist,
11    MemberlistProvider,
12    SegmentDirectory,
13)
14from chromadb.telemetry.opentelemetry import (
15    OpenTelemetryGranularity,
16    add_attributes_to_current_span,
17    trace_method,
18)
19from chromadb.types import Segment
20from chromadb.utils.rendezvous_hash import assign, murmur3hasher
21
22# These could go in config but given that they will rarely change, they are here for now to avoid
23# polluting the config file further.
24WATCH_TIMEOUT_SECONDS = 60
25KUBERNETES_NAMESPACE = "chroma"
26KUBERNETES_GROUP = "chroma.cluster"
27HEADLESS_SERVICE = "svc.cluster.local"
28
29
30class MockMemberlistProvider(MemberlistProvider, EnforceOverrides):
31    """A mock memberlist provider for testing"""
32
33    _memberlist: Memberlist
34
35    def __init__(self, system: System):
36        super().__init__(system)
37        self._memberlist = [
38            Member(id="a", ip="10.0.0.1", node="node1"),
39            Member(id="b", ip="10.0.0.2", node="node2"),
40            Member(id="c", ip="10.0.0.3", node="node3"),
41        ]
42
43    @override
44    def get_memberlist(self) -> Memberlist:
45        return self._memberlist
46
47    @override
48    def set_memberlist_name(self, memberlist: str) -> None:
49        pass  # The mock provider does not need to set the memberlist name
50
51    def update_memberlist(self, memberlist: Memberlist) -> None:
52        """Updates the memberlist and calls all registered callbacks. This mocks an update from a k8s CR"""
53        self._memberlist = memberlist
54        for callback in self.callbacks:
55            callback(memberlist)
56
57
58class CustomResourceMemberlistProvider(MemberlistProvider, EnforceOverrides):
59    """A memberlist provider that uses a k8s custom resource to store the memberlist"""
60
61    _kubernetes_api: client.CustomObjectsApi
62    _memberlist_name: Optional[str]
63    _curr_memberlist: Optional[Memberlist]
64    _curr_memberlist_mutex: threading.Lock
65    _watch_thread: Optional[threading.Thread]
66    _kill_watch_thread: threading.Event
67    _done_waiting_for_reset: threading.Event
68
69    def __init__(self, system: System):
70        super().__init__(system)
71        config.load_config()
72        self._kubernetes_api = client.CustomObjectsApi()
73        self._watch_thread = None
74        self._memberlist_name = None
75        self._curr_memberlist = None
76        self._curr_memberlist_mutex = threading.Lock()
77        self._kill_watch_thread = threading.Event()
78        self._done_waiting_for_reset = threading.Event()
79
80    @override
81    def start(self) -> None:
82        if self._memberlist_name is None:
83            raise ValueError("Memberlist name must be set before starting")
84        self.get_memberlist()
85        self._done_waiting_for_reset.clear()
86        self._watch_worker_memberlist()
87        return super().start()
88
89    @override
90    def stop(self) -> None:
91        self._curr_memberlist = None
92        self._memberlist_name = None
93
94        # Stop the watch thread
95        self._kill_watch_thread.set()
96        if self._watch_thread is not None:
97            self._watch_thread.join()
98        self._watch_thread = None
99        self._kill_watch_thread.clear()
100        self._done_waiting_for_reset.clear()
101        return super().stop()
102
103    @override
104    def reset_state(self) -> None:
105        # Reset the memberlist in kubernetes, and wait for it to
106        # get propagated back again
107        # Note that the component must be running in order to reset the state
108        if not self._system.settings.require("allow_reset"):
109            raise ValueError(
110                "Resetting the database is not allowed. Set `allow_reset` to true in the config in tests or other non-production environments where reset should be permitted."
111            )
112        if self._memberlist_name:
113            self._done_waiting_for_reset.clear()
114            self._kubernetes_api.patch_namespaced_custom_object(
115                group=KUBERNETES_GROUP,
116                version="v1",
117                namespace=KUBERNETES_NAMESPACE,
118                plural="memberlists",
119                name=self._memberlist_name,
120                body={
121                    "kind": "MemberList",
122                    "spec": {"members": []},
123                },
124            )
125            self._done_waiting_for_reset.wait(5.0)
126            # TODO: For some reason the above can flake and the memberlist won't be populated
127            # Given that this is a test harness, just sleep for an additional 500ms for now
128            # We should understand why this flaps
129            time.sleep(0.5)
130
131    @override
132    def get_memberlist(self) -> Memberlist:
133        if self._curr_memberlist is None:
134            self._curr_memberlist = self._fetch_memberlist()
135        return self._curr_memberlist
136
137    @override
138    def set_memberlist_name(self, memberlist: str) -> None:
139        self._memberlist_name = memberlist
140
141    def _fetch_memberlist(self) -> Memberlist:
142        api_response = self._kubernetes_api.get_namespaced_custom_object(
143            group=KUBERNETES_GROUP,
144            version="v1",
145            namespace=KUBERNETES_NAMESPACE,
146            plural="memberlists",
147            name=f"{self._memberlist_name}",
148        )
149        api_response = cast(Dict[str, Any], api_response)
150        if "spec" not in api_response:
151            return []
152        response_spec = cast(Dict[str, Any], api_response["spec"])
153        return self._parse_response_memberlist(response_spec)
154
155    def _watch_worker_memberlist(self) -> None:
156        # TODO: We may want to make this watch function a library function that can be used by other
157        # components that need to watch k8s custom resources.
158        def run_watch() -> None:
159            w = watch.Watch()
160
161            def do_watch() -> None:
162                for event in w.stream(
163                    self._kubernetes_api.list_namespaced_custom_object,
164                    group=KUBERNETES_GROUP,
165                    version="v1",
166                    namespace=KUBERNETES_NAMESPACE,
167                    plural="memberlists",
168                    field_selector=f"metadata.name={self._memberlist_name}",
169                    timeout_seconds=WATCH_TIMEOUT_SECONDS,
170                ):
171                    event = cast(Dict[str, Any], event)
172                    response_spec = event["object"]["spec"]
173                    response_spec = cast(Dict[str, Any], response_spec)
174                    with self._curr_memberlist_mutex:
175                        self._curr_memberlist = self._parse_response_memberlist(
176                            response_spec
177                        )
178                    self._notify(self._curr_memberlist)
179                    if (
180                        self._system.settings.require("allow_reset")
181                        and not self._done_waiting_for_reset.is_set()
182                        and len(self._curr_memberlist) > 0
183                    ):
184                        self._done_waiting_for_reset.set()
185
186            # Watch the custom resource for changes
187            # Watch with a timeout and retry so we can gracefully stop this if needed
188            while not self._kill_watch_thread.is_set():
189                try:
190                    do_watch()
191                except ApiException as e:
192                    # If status code is 410, the watch has expired and we need to start a new one.
193                    if e.status == 410:
194                        pass
195            return
196
197        if self._watch_thread is None:
198            thread = threading.Thread(target=run_watch, daemon=True)
199            thread.start()
200            self._watch_thread = thread
201        else:
202            raise Exception("A watch thread is already running.")
203
204    def _parse_response_memberlist(
205        self, api_response_spec: Dict[str, Any]
206    ) -> Memberlist:
207        if "members" not in api_response_spec:
208            return []
209        parsed = []
210        for m in api_response_spec["members"]:
211            id = m["member_id"]
212            ip = m["member_ip"] if "member_ip" in m else ""
213            node = m["member_node_name"] if "member_node_name" in m else ""
214            parsed.append(Member(id=id, ip=ip, node=node))
215        return parsed
216
217    def _notify(self, memberlist: Memberlist) -> None:
218        for callback in self.callbacks:
219            callback(memberlist)
220
221
222class RendezvousHashSegmentDirectory(SegmentDirectory, EnforceOverrides):
223    _memberlist_provider: MemberlistProvider
224    _curr_memberlist_mutex: threading.Lock
225    _curr_memberlist: Optional[Memberlist]
226    _routing_mode: RoutingMode
227
228    def __init__(self, system: System):
229        super().__init__(system)
230        self._memberlist_provider = self.require(MemberlistProvider)
231        memberlist_name = system.settings.require("worker_memberlist_name")
232        self._memberlist_provider.set_memberlist_name(memberlist_name)
233        self._routing_mode = system.settings.require(
234            "chroma_segment_directory_routing_mode"
235        )
236
237        self._curr_memberlist = None
238        self._curr_memberlist_mutex = threading.Lock()
239
240    @override
241    def start(self) -> None:
242        self._curr_memberlist = self._memberlist_provider.get_memberlist()
243        self._memberlist_provider.register_updated_memberlist_callback(
244            self._update_memberlist
245        )
246        return super().start()
247
248    @override
249    def stop(self) -> None:
250        self._memberlist_provider.unregister_updated_memberlist_callback(
251            self._update_memberlist
252        )
253        return super().stop()
254
255    @override
256    def get_segment_endpoints(self, segment: Segment, n: int) -> List[str]:
257        if self._curr_memberlist is None or len(self._curr_memberlist) == 0:
258            raise ValueError("Memberlist is not initialized")
259
260        # assign() will throw an error if n is greater than the number of members
261        # clamp n to the number of members to align with the contract of this method
262        # which is to return at most n endpoints
263        n = min(n, len(self._curr_memberlist))
264
265        # Check if all members in the memberlist have a node set,
266        # if so, route using the node
267
268        # NOTE(@hammadb) 1/8/2024: This is to handle the migration between routing
269        # using the member id and routing using the node name
270        # We want to route using the node name over the member id
271        # because the node may have a disk cache that we want a
272        # stable identifier for over deploys.
273        can_use_node_routing = (
274            all([m.node != "" and len(m.node) != 0 for m in self._curr_memberlist])
275            and self._routing_mode == RoutingMode.NODE
276        )
277        if can_use_node_routing:
278            # If we are using node routing and the segments
279            assignments = assign(
280                segment["collection"].hex,
281                [m.node for m in self._curr_memberlist],
282                murmur3hasher,
283                n,
284            )
285        else:
286            # Query to the same collection should end up on the same endpoint
287            assignments = assign(
288                segment["collection"].hex,
289                [m.id for m in self._curr_memberlist],
290                murmur3hasher,
291                n,
292            )
293        assignments_set = set(assignments)
294        out_endpoints = []
295        for member in self._curr_memberlist:
296            is_chosen_with_node_routing = (
297                can_use_node_routing and member.node in assignments_set
298            )
299            is_chosen_with_id_routing = (
300                not can_use_node_routing and member.id in assignments_set
301            )
302            if is_chosen_with_node_routing or is_chosen_with_id_routing:
303                # If the memberlist has an ip, use it, otherwise use the member id with the headless service
304                # this is for backwards compatibility with the old memberlist which only had ids
305                if member.ip is not None and member.ip != "":
306                    endpoint = f"{member.ip}:50051"
307                    out_endpoints.append(endpoint)
308                else:
309                    service_name = self.extract_service_name(member.id)
310                    endpoint = f"{member.id}.{service_name}.{KUBERNETES_NAMESPACE}.{HEADLESS_SERVICE}:50051"
311                    out_endpoints.append(endpoint)
312        return out_endpoints
313
314    @override
315    def register_updated_segment_callback(
316        self, callback: Callable[[Segment], None]
317    ) -> None:
318        raise NotImplementedError()
319
320    @trace_method(
321        "RendezvousHashSegmentDirectory._update_memberlist",
322        OpenTelemetryGranularity.ALL,
323    )
324    def _update_memberlist(self, memberlist: Memberlist) -> None:
325        with self._curr_memberlist_mutex:
326            add_attributes_to_current_span(
327                {"new_memberlist": [m.id for m in memberlist]}
328            )
329            self._curr_memberlist = memberlist
330
331    def extract_service_name(self, pod_name: str) -> Optional[str]:
332        # Split the pod name by the hyphen
333        parts = pod_name.split("-")
334        # The service name is expected to be the prefix before the last hyphen
335        if len(parts) > 1:
336            return "-".join(parts[:-1])
337        return None
338 
codekingpro/portable-devtools · Team Ai