codekingpro/portable-devtools
115k
1from typing import ClassVar, Dict, Optional
2import logging
3import threading
4import uuid
5from chromadb.api import ServerAPI
6from chromadb.api.base_http_client import BaseHTTPClient
7from chromadb.config import Settings, System
8from chromadb.telemetry.product import ProductTelemetryClient
9from chromadb.telemetry.product.events import ClientStartEvent
10
11logger = logging.getLogger(__name__)
12
13
14class SharedSystemClient:
15 _identifier_to_system: ClassVar[Dict[str, System]] = {}
16 _identifier_to_refcount: ClassVar[Dict[str, int]] = {}
17 _refcount_lock: ClassVar[threading.Lock] = threading.Lock()
18 _identifier: str
19
20 def __init__(
21 self,
22 settings: Settings = Settings(),
23 ) -> None:
24 self._identifier = SharedSystemClient._get_identifier_from_settings(settings)
25 SharedSystemClient._create_system_if_not_exists(self._identifier, settings)
26 SharedSystemClient._increment_refcount(self._identifier)
27
28 @classmethod
29 def _create_system_if_not_exists(
30 cls, identifier: str, settings: Settings
31 ) -> System:
32 if identifier not in cls._identifier_to_system:
33 new_system = System(settings)
34 cls._identifier_to_system[identifier] = new_system
35
36 new_system.instance(ProductTelemetryClient)
37 new_system.instance(ServerAPI)
38
39 new_system.start()
40 else:
41 previous_system = cls._identifier_to_system[identifier]
42
43 # For now, the settings must match
44 if previous_system.settings != settings:
45 raise ValueError(
46 f"An instance of Chroma already exists for {identifier} with different settings"
47 )
48
49 return cls._identifier_to_system[identifier]
50
51 @staticmethod
52 def _get_identifier_from_settings(settings: Settings) -> str:
53 identifier = ""
54 api_impl = settings.chroma_api_impl
55
56 if api_impl is None:
57 raise ValueError("Chroma API implementation must be set in settings")
58 elif api_impl in [
59 "chromadb.api.segment.SegmentAPI",
60 "chromadb.api.rust.RustBindingsAPI",
61 ]:
62 if settings.is_persistent:
63 identifier = settings.persist_directory
64 else:
65 identifier = (
66 "ephemeral" # TODO: support pathing and multiple ephemeral clients
67 )
68 elif api_impl in [
69 "chromadb.api.fastapi.FastAPI",
70 "chromadb.api.async_fastapi.AsyncFastAPI",
71 ]:
72 # FastAPI clients can all use unique system identifiers since their configurations can be independent, e.g. different auth tokens
73 identifier = str(uuid.uuid4())
74 else:
75 raise ValueError(f"Unsupported Chroma API implementation {api_impl}")
76
77 return identifier
78
79 @staticmethod
80 def _populate_data_from_system(system: System) -> str:
81 identifier = SharedSystemClient._get_identifier_from_settings(system.settings)
82 SharedSystemClient._identifier_to_system[identifier] = system
83 return identifier
84
85 @classmethod
86 def from_system(cls, system: System) -> "SharedSystemClient":
87 """Create a client from an existing system. This is useful for testing and debugging."""
88
89 SharedSystemClient._populate_data_from_system(system)
90 instance = cls(system.settings)
91 return instance
92
93 @classmethod
94 def _increment_refcount(cls, identifier: str) -> None:
95 """Increment the reference count for a system identifier."""
96 with cls._refcount_lock:
97 if identifier not in cls._identifier_to_refcount:
98 cls._identifier_to_refcount[identifier] = 0
99 cls._identifier_to_refcount[identifier] += 1
100
101 @classmethod
102 def _decrement_refcount(cls, identifier: str) -> int:
103 """Decrement the reference count for a system identifier and return the new count."""
104 with cls._refcount_lock:
105 if identifier in cls._identifier_to_refcount:
106 cls._identifier_to_refcount[identifier] -= 1
107 count = cls._identifier_to_refcount[identifier]
108 if count <= 0:
109 del cls._identifier_to_refcount[identifier]
110 return count
111 return 0
112
113 @classmethod
114 def _release_system(cls, identifier: str) -> None:
115 """Decrement refcount and stop the system if this was the last reference.
116
117 This consolidates the "decrement + conditional stop" pattern used in
118 both Client.close() and the Client.__init__ exception handler.
119 """
120 refcount = cls._decrement_refcount(identifier)
121 if refcount <= 0:
122 system = cls._identifier_to_system.pop(identifier, None)
123 if system is not None:
124 system.stop()
125
126 @staticmethod
127 def clear_system_cache() -> None:
128 SharedSystemClient._identifier_to_system = {}
129 SharedSystemClient._identifier_to_refcount = {}
130
131 @property
132 def _system(self) -> System:
133 return SharedSystemClient._identifier_to_system[self._identifier]
134
135 def _submit_client_start_event(self) -> None:
136 telemetry_client = self._system.instance(ProductTelemetryClient)
137 telemetry_client.capture(ClientStartEvent())
138
139 @staticmethod
140 def get_chroma_cloud_api_key_from_clients() -> Optional[str]:
141 """
142 Try to extract api key from existing client instances by checking httpx session headers.
143
144 Requirements to pull api key:
145 - must be a BaseHTTPClient instance (ignore RustBindingsAPI and SegmentAPI)
146 - must have "api.trychroma.com" or "gcp.trychroma.com" in the _api_url (ignore local/self-hosted instances)
147 - must have "x-chroma-token" or "X-Chroma-Token" in the headers
148
149 Returns:
150 The first api key found, or None if no client instances have api keys set.
151 """
152
153 api_keys: list[str] = []
154 systems_snapshot = list(SharedSystemClient._identifier_to_system.values())
155 for system in systems_snapshot:
156 try:
157 server_api = system.instance(ServerAPI)
158
159 if not isinstance(server_api, BaseHTTPClient):
160 # RustBindingsAPI and SegmentAPI don't have HTTP headers
161 continue
162
163 # Only pull api key if the url contains the chroma cloud url
164 api_url = server_api.get_api_url()
165 if (
166 "api.trychroma.com" not in api_url
167 and "gcp.trychroma.com" not in api_url
168 ):
169 continue
170
171 headers = server_api.get_request_headers()
172 api_key = None
173 for key, value in headers.items():
174 if key.lower() == "x-chroma-token":
175 api_key = value
176 break
177
178 if api_key:
179 api_keys.append(api_key)
180 except Exception:
181 # If we can't access the ServerAPI instance, continue to the next
182 continue
183
184 if not api_keys:
185 return None
186
187 # log if multiple viable api keys found
188 if len(api_keys) > 1:
189 logger.info(
190 f"Multiple Chroma Cloud clients found, using API key starting with {api_keys[0][:8]}..."
191 )
192
193 return api_keys[0]
194 