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