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

1import logging 

2from typing import Optional 

3 

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 

28 

29log = logging.getLogger(__name__) 

30 

31 

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 

44 

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 ) 

62 

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 

69 

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) 

73 

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 ) 

90 

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]] 

96 

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 

108 

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 ) 

118 

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 

129 

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 

143 

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 ) 

149 

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] 

154 

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) 

163 

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 ) 

169 

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] 

174 

175 collection.upsert(ids=ids, documents=documents, embeddings=embeddings, metadatas=metadatas) 

176 

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 

195 

196 def reset(self): 

197 # Resets the database. This will delete all collections and item entries. 

198 return self.client.reset()