Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
local.py270 linesDownload Raw Back to manager
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 
codekingpro/portable-devtools · Team Ai