Coverage for open_webui/retrieval/vector/dbs/chroma.py: 51%
87 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import logging
2from typing import Optional
4import chromadb
5from chromadb import Settings
6from chromadb.errors import NotFoundError
7from chromadb.utils.batch_utils import create_batches
8from open_webui.config import (
9 CHROMA_CLIENT_AUTH_CREDENTIALS,
10 CHROMA_CLIENT_AUTH_PROVIDER,
11 CHROMA_DATA_PATH,
12 CHROMA_DATABASE,
13 CHROMA_HTTP_HEADERS,
14 CHROMA_HTTP_HOST,
15 CHROMA_HTTP_PORT,
16 CHROMA_HTTP_SSL,
17 CHROMA_TENANT,
18)
19from open_webui.env import USE_SLIM
20from fastapi import HTTPException
21from open_webui.retrieval.vector.main import (
22 GetResult,
23 SearchResult,
24 VectorDBBase,
25 VectorItem,
26)
27from open_webui.retrieval.vector.utils import process_metadata
29log = logging.getLogger(__name__)
32class ChromaClient(VectorDBBase):
33 def __init__(self):
34 if USE_SLIM and not CHROMA_HTTP_HOST: 34 ↛ 35line 34 didn't jump to line 35 because the condition on line 34 was never true
35 raise HTTPException(503, 'Configure CHROMA_HTTP_HOST: embedded Chroma is unavailable in slim.')
36 settings_dict = {
37 'allow_reset': True,
38 'anonymized_telemetry': False,
39 }
40 if CHROMA_CLIENT_AUTH_PROVIDER is not None: 40 ↛ 42line 40 didn't jump to line 42 because the condition on line 40 was always true
41 settings_dict['chroma_client_auth_provider'] = CHROMA_CLIENT_AUTH_PROVIDER
42 if CHROMA_CLIENT_AUTH_CREDENTIALS is not None: 42 ↛ 45line 42 didn't jump to line 45 because the condition on line 42 was always true
43 settings_dict['chroma_client_auth_credentials'] = CHROMA_CLIENT_AUTH_CREDENTIALS
45 if CHROMA_HTTP_HOST != '': 45 ↛ 46line 45 didn't jump to line 46 because the condition on line 45 was never true
46 self.client = chromadb.HttpClient(
47 host=CHROMA_HTTP_HOST,
48 port=CHROMA_HTTP_PORT,
49 headers=CHROMA_HTTP_HEADERS,
50 ssl=CHROMA_HTTP_SSL,
51 tenant=CHROMA_TENANT,
52 database=CHROMA_DATABASE,
53 settings=Settings(**settings_dict),
54 )
55 else:
56 self.client = chromadb.PersistentClient(
57 path=CHROMA_DATA_PATH,
58 settings=Settings(**settings_dict),
59 tenant=CHROMA_TENANT,
60 database=CHROMA_DATABASE,
61 )
63 def has_collection(self, collection_name: str) -> bool:
64 try:
65 self.client.get_collection(name=collection_name, embedding_function=None)
66 return True
67 except NotFoundError:
68 return False
70 def delete_collection(self, collection_name: str):
71 # Delete the collection based on the collection name.
72 return self.client.delete_collection(name=collection_name)
74 def search(
75 self,
76 collection_name: str,
77 vectors: list[list[float | int]],
78 filter: Optional[dict] = None,
79 limit: int = 10,
80 ) -> Optional[SearchResult]:
81 # Search for the nearest neighbor items based on the vectors and return 'limit' number of results.
82 try:
83 collection = self.client.get_collection(name=collection_name, embedding_function=None)
84 if collection:
85 result = collection.query(
86 query_embeddings=vectors,
87 n_results=limit,
88 where=filter,
89 )
91 # chromadb has cosine distance, 2 (worst) -> 0 (best). Re-odering to 0 -> 1
92 # https://docs.trychroma.com/docs/collections/configure cosine equation
93 distances: list = result['distances'][0]
94 distances = [2 - dist for dist in distances]
95 distances = [[dist / 2 for dist in distances]]
97 return SearchResult(
98 **{
99 'ids': result['ids'],
100 'distances': distances,
101 'documents': result['documents'],
102 'metadatas': result['metadatas'],
103 }
104 )
105 return None
106 except Exception as e:
107 return None
109 def query(self, collection_name: str, filter: dict, limit: Optional[int] = None) -> Optional[GetResult]:
110 # Query the items from the collection based on the filter.
111 try:
112 collection = self.client.get_collection(name=collection_name, embedding_function=None)
113 if collection:
114 result = collection.get(
115 where=filter,
116 limit=limit,
117 )
119 return GetResult(
120 **{
121 'ids': [result['ids']],
122 'documents': [result['documents']],
123 'metadatas': [result['metadatas']],
124 }
125 )
126 return None
127 except Exception:
128 return None
130 def get(self, collection_name: str) -> Optional[GetResult]:
131 # Get all the items in the collection.
132 collection = self.client.get_collection(name=collection_name, embedding_function=None)
133 if collection:
134 result = collection.get()
135 return GetResult(
136 **{
137 'ids': [result['ids']],
138 'documents': [result['documents']],
139 'metadatas': [result['metadatas']],
140 }
141 )
142 return None
144 def insert(self, collection_name: str, items: list[VectorItem]):
145 # Insert the items into the collection, if the collection does not exist, it will be created.
146 collection = self.client.get_or_create_collection(
147 name=collection_name, metadata={'hnsw:space': 'cosine'}, embedding_function=None
148 )
150 ids = [item['id'] for item in items]
151 documents = [item['text'] for item in items]
152 embeddings = [item['vector'] for item in items]
153 metadatas = [process_metadata(item['metadata']) for item in items]
155 for batch in create_batches(
156 api=self.client,
157 documents=documents,
158 embeddings=embeddings,
159 ids=ids,
160 metadatas=metadatas,
161 ):
162 collection.add(*batch)
164 def upsert(self, collection_name: str, items: list[VectorItem]):
165 # Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created.
166 collection = self.client.get_or_create_collection(
167 name=collection_name, metadata={'hnsw:space': 'cosine'}, embedding_function=None
168 )
170 ids = [item['id'] for item in items]
171 documents = [item['text'] for item in items]
172 embeddings = [item['vector'] for item in items]
173 metadatas = [process_metadata(item['metadata']) for item in items]
175 collection.upsert(ids=ids, documents=documents, embeddings=embeddings, metadatas=metadatas)
177 def delete(
178 self,
179 collection_name: str,
180 ids: Optional[list[str]] = None,
181 filter: Optional[dict] = None,
182 ):
183 # Delete the items from the collection based on the ids.
184 try:
185 collection = self.client.get_collection(name=collection_name, embedding_function=None)
186 if collection:
187 if ids:
188 collection.delete(ids=ids)
189 elif filter:
190 collection.delete(where=filter)
191 except Exception as e:
192 # If collection doesn't exist, that's fine - nothing to delete
193 log.debug('Attempted to delete from non-existent collection %s. Ignoring.', collection_name)
194 pass
196 def reset(self):
197 # Resets the database. This will delete all collections and item entries.
198 return self.client.reset()