codekingpro/portable-devtools
114k
1from threading import Lock
2from chromadb.segment import (
3 SegmentImplementation,
4 SegmentManager,
5 MetadataReader,
6 SegmentType,
7 VectorReader,
8 S,
9)
10import logging
11from chromadb.segment.impl.manager.cache.cache import (
12 SegmentLRUCache,
13 BasicCache,
14 SegmentCache,
15)
16import os
17
18from chromadb.config import System, get_class
19from chromadb.db.system import SysDB
20from overrides import override
21from chromadb.segment.impl.vector.local_persistent_hnsw import (
22 PersistentLocalHnswSegment,
23)
24from chromadb.telemetry.opentelemetry import (
25 OpenTelemetryClient,
26 OpenTelemetryGranularity,
27 trace_method,
28)
29from chromadb.types import Collection, Operation, Segment, SegmentScope, Metadata
30from typing import Dict, Type, Sequence, Optional, cast
31from uuid import UUID, uuid4
32import platform
33
34from chromadb.utils.lru_cache import LRUCache
35from chromadb.utils.directory import get_directory_size
36
37
38if platform.system() != "Windows":
39 import resource
40elif platform.system() == "Windows":
41 import ctypes
42
43SEGMENT_TYPE_IMPLS = {
44 SegmentType.SQLITE: "chromadb.segment.impl.metadata.sqlite.SqliteMetadataSegment",
45 SegmentType.HNSW_LOCAL_MEMORY: "chromadb.segment.impl.vector.local_hnsw.LocalHnswSegment",
46 SegmentType.HNSW_LOCAL_PERSISTED: "chromadb.segment.impl.vector.local_persistent_hnsw.PersistentLocalHnswSegment",
47}
48
49
50class LocalSegmentManager(SegmentManager):
51 _sysdb: SysDB
52 _system: System
53 _opentelemetry_client: OpenTelemetryClient
54 _instances: Dict[UUID, SegmentImplementation]
55 _vector_instances_file_handle_cache: LRUCache[
56 UUID, PersistentLocalHnswSegment
57 ] # LRU cache to manage file handles across vector segment instances
58 _vector_segment_type: SegmentType = SegmentType.HNSW_LOCAL_MEMORY
59 _lock: Lock
60 _max_file_handles: int
61
62 def __init__(self, system: System):
63 super().__init__(system)
64 self._sysdb = self.require(SysDB)
65 self._system = system
66 self._opentelemetry_client = system.require(OpenTelemetryClient)
67 self.logger = logging.getLogger(__name__)
68 self._instances = {}
69 self.segment_cache: Dict[SegmentScope, SegmentCache] = {
70 SegmentScope.METADATA: BasicCache() # type: ignore[no-untyped-call]
71 }
72 if (
73 system.settings.chroma_segment_cache_policy == "LRU"
74 and system.settings.chroma_memory_limit_bytes > 0
75 ):
76 self.segment_cache[SegmentScope.VECTOR] = SegmentLRUCache(
77 capacity=system.settings.chroma_memory_limit_bytes,
78 callback=lambda k, v: self.callback_cache_evict(v),
79 size_func=lambda k: self._get_segment_disk_size(k),
80 )
81 else:
82 self.segment_cache[SegmentScope.VECTOR] = BasicCache() # type: ignore[no-untyped-call]
83
84 self._lock = Lock()
85
86 # TODO: prototyping with distributed segment for now, but this should be a configurable option
87 # we need to think about how to handle this configuration
88 if self._system.settings.require("is_persistent"):
89 self._vector_segment_type = SegmentType.HNSW_LOCAL_PERSISTED
90 if platform.system() != "Windows":
91 self._max_file_handles = resource.getrlimit(resource.RLIMIT_NOFILE)[0]
92 else:
93 self._max_file_handles = ctypes.windll.msvcrt._getmaxstdio() # type: ignore
94 segment_limit = (
95 self._max_file_handles
96 # This is integer division in Python 3, and not a comment.
97 // PersistentLocalHnswSegment.get_file_handle_count()
98 )
99 self._vector_instances_file_handle_cache = LRUCache(
100 segment_limit, callback=lambda _, v: v.close_persistent_index()
101 )
102
103 @trace_method(
104 "LocalSegmentManager.callback_cache_evict",
105 OpenTelemetryGranularity.OPERATION_AND_SEGMENT,
106 )
107 def callback_cache_evict(self, segment: Segment) -> None:
108 collection_id = segment["collection"]
109 self.logger.info(f"LRU cache evict collection {collection_id}")
110 instance = self._instance(segment)
111 instance.stop()
112 del self._instances[segment["id"]]
113
114 @override
115 def start(self) -> None:
116 for instance in self._instances.values():
117 instance.start()
118 super().start()
119
120 @override
121 def stop(self) -> None:
122 for instance in self._instances.values():
123 instance.stop()
124 super().stop()
125
126 @override
127 def reset_state(self) -> None:
128 for instance in self._instances.values():
129 instance.stop()
130 instance.reset_state()
131 self._instances = {}
132 self.segment_cache[SegmentScope.VECTOR].reset()
133 super().reset_state()
134
135 @trace_method(
136 "LocalSegmentManager.prepare_segments_for_new_collection",
137 OpenTelemetryGranularity.OPERATION_AND_SEGMENT,
138 )
139 @override
140 def prepare_segments_for_new_collection(
141 self, collection: Collection
142 ) -> Sequence[Segment]:
143 vector_segment = _segment(
144 self._vector_segment_type, SegmentScope.VECTOR, collection
145 )
146 metadata_segment = _segment(
147 SegmentType.SQLITE, SegmentScope.METADATA, collection
148 )
149 return [vector_segment, metadata_segment]
150
151 @trace_method(
152 "LocalSegmentManager.delete_segments",
153 OpenTelemetryGranularity.OPERATION_AND_SEGMENT,
154 )
155 @override
156 def delete_segments(self, collection_id: UUID) -> Sequence[UUID]:
157 segments = self._sysdb.get_segments(collection=collection_id)
158 for segment in segments:
159 if segment["id"] in self._instances:
160 if segment["type"] == SegmentType.HNSW_LOCAL_PERSISTED.value:
161 instance = self.get_segment(collection_id, VectorReader)
162 instance.delete()
163 elif segment["type"] == SegmentType.SQLITE.value:
164 instance = self.get_segment(collection_id, MetadataReader) # type: ignore[assignment]
165 instance.delete()
166 del self._instances[segment["id"]]
167 if segment["scope"] is SegmentScope.VECTOR:
168 self.segment_cache[SegmentScope.VECTOR].pop(collection_id)
169 if segment["scope"] is SegmentScope.METADATA:
170 self.segment_cache[SegmentScope.METADATA].pop(collection_id)
171 return [s["id"] for s in segments]
172
173 def _get_segment_disk_size(self, collection_id: UUID) -> int:
174 segments = self._sysdb.get_segments(
175 collection=collection_id, scope=SegmentScope.VECTOR
176 )
177 if len(segments) == 0:
178 return 0
179 # With local segment manager (single server chroma), a collection always have one segment.
180 size = get_directory_size(
181 os.path.join(
182 self._system.settings.require("persist_directory"),
183 str(segments[0]["id"]),
184 )
185 )
186 return size
187
188 @trace_method(
189 "LocalSegmentManager._get_segment_sysdb",
190 OpenTelemetryGranularity.OPERATION_AND_SEGMENT,
191 )
192 def _get_segment_sysdb(self, collection_id: UUID, scope: SegmentScope) -> Segment:
193 segments = self._sysdb.get_segments(collection=collection_id, scope=scope)
194 known_types = set([k.value for k in SEGMENT_TYPE_IMPLS.keys()])
195 # Get the first segment of a known type
196 segment = next(filter(lambda s: s["type"] in known_types, segments))
197 return segment
198
199 @trace_method(
200 "LocalSegmentManager.get_segment",
201 OpenTelemetryGranularity.OPERATION_AND_SEGMENT,
202 )
203 def get_segment(self, collection_id: UUID, type: Type[S]) -> S:
204 if type == MetadataReader:
205 scope = SegmentScope.METADATA
206 elif type == VectorReader:
207 scope = SegmentScope.VECTOR
208 else:
209 raise ValueError(f"Invalid segment type: {type}")
210
211 segment = self.segment_cache[scope].get(collection_id)
212 if segment is None:
213 segment = self._get_segment_sysdb(collection_id, scope)
214 self.segment_cache[scope].set(collection_id, segment)
215
216 # Instances must be atomically created, so we use a lock to ensure that only one thread
217 # creates the instance.
218 with self._lock:
219 instance = self._instance(segment)
220 return cast(S, instance)
221
222 @trace_method(
223 "LocalSegmentManager.hint_use_collection",
224 OpenTelemetryGranularity.OPERATION_AND_SEGMENT,
225 )
226 @override
227 def hint_use_collection(self, collection_id: UUID, hint_type: Operation) -> None:
228 # The local segment manager responds to hints by pre-loading both the metadata and vector
229 # segments for the given collection.
230 for type in [MetadataReader, VectorReader]:
231 # Just use get_segment to load the segment into the cache
232 instance = self.get_segment(collection_id, type)
233 # If the segment is a vector segment, we need to keep segments in an LRU cache
234 # to avoid hitting the OS file handle limit.
235 if type == VectorReader and self._system.settings.require("is_persistent"):
236 instance = cast(PersistentLocalHnswSegment, instance)
237 instance.open_persistent_index()
238 self._vector_instances_file_handle_cache.set(collection_id, instance)
239
240 def _cls(self, segment: Segment) -> Type[SegmentImplementation]:
241 classname = SEGMENT_TYPE_IMPLS[SegmentType(segment["type"])]
242 cls = get_class(classname, SegmentImplementation)
243 return cls
244
245 def _instance(self, segment: Segment) -> SegmentImplementation:
246 if segment["id"] not in self._instances:
247 cls = self._cls(segment)
248 instance = cls(self._system, segment)
249 instance.start()
250 self._instances[segment["id"]] = instance
251 return self._instances[segment["id"]]
252
253
254def _segment(type: SegmentType, scope: SegmentScope, collection: Collection) -> Segment:
255 """Create a metadata dict, propagating metadata correctly for the given segment type."""
256 cls = get_class(SEGMENT_TYPE_IMPLS[type], SegmentImplementation)
257 collection_metadata = collection.metadata
258 metadata: Optional[Metadata] = None
259 if collection_metadata:
260 metadata = cls.propagate_collection_metadata(collection_metadata)
261
262 return Segment(
263 id=uuid4(),
264 type=type.value,
265 scope=scope,
266 collection=collection.id,
267 metadata=metadata,
268 file_paths={},
269 )
270 