Coverage for open_webui/retrieval/external.py: 9%

182 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1import asyncio 

2import logging 

3import re 

4import time 

5from typing import Any, Optional 

6 

7from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX 

8from open_webui.models.config import Config 

9from open_webui.models.knowledge import KnowledgeModel 

10 

11log = logging.getLogger(__name__) 

12 

13EXTERNAL_KNOWLEDGE_CONNECTIONS_CONFIG_KEY = 'external_knowledge.connections' 

14IDENTIFIER_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]*$') 

15 

16 

17async def _get_external_connection(connection_id: str) -> Optional[dict]: 

18 connections = await Config.get(EXTERNAL_KNOWLEDGE_CONNECTIONS_CONFIG_KEY, []) or [] 

19 return next((connection for connection in connections if connection.get('id') == connection_id), None) 

20 

21 

22def _get_path(data: Any, path: Optional[str], default=None): 

23 if not path: 

24 return default 

25 value = data 

26 for part in path.split('.'): 

27 if isinstance(value, dict): 

28 value = value.get(part, default) 

29 else: 

30 return default 

31 return value 

32 

33 

34def _normalize_result(result: dict, mapping: dict, knowledge: KnowledgeModel, distance: Optional[float] = None) -> dict: 

35 content = _get_path(result, mapping.get('content_field', 'content'), '') 

36 title = _get_path(result, mapping.get('title_field', 'title'), None) 

37 source = _get_path(result, mapping.get('source_field', 'source'), None) 

38 url = _get_path(result, mapping.get('url_field', 'url'), None) 

39 document_id = _get_path(result, mapping.get('document_id_field', 'document_id'), None) 

40 page = _get_path(result, mapping.get('page_field', 'page'), None) 

41 metadata = _get_path(result, mapping.get('metadata_field', 'metadata'), {}) or {} 

42 score = _get_path(result, mapping.get('score_field', 'score'), distance) 

43 

44 if not isinstance(metadata, dict): 

45 metadata = {'external_metadata': metadata} 

46 

47 source_name = source or title or metadata.get('source') or metadata.get('name') or knowledge.name 

48 metadata.update( 

49 { 

50 'name': title or source_name, 

51 'source': source_name, 

52 'url': url, 

53 'file_id': document_id or f'external-{knowledge.id}', 

54 'knowledge_id': knowledge.id, 

55 'knowledge_name': knowledge.name, 

56 'external': True, 

57 } 

58 ) 

59 if page is not None: 

60 metadata['page'] = page 

61 if document_id is not None: 

62 metadata['document_id'] = document_id 

63 

64 return { 

65 'content': content, 

66 'metadata': metadata, 

67 'distance': score, 

68 } 

69 

70 

71def _source_config(knowledge: KnowledgeModel) -> dict: 

72 external = (knowledge.meta or {}).get('external', {}) 

73 source = external.get('source') or {} 

74 return source.get('config') or {} 

75 

76 

77def _root_field(path: Optional[str]) -> Optional[str]: 

78 if not path: 

79 return None 

80 return path.split('.')[0] 

81 

82 

83def _safe_identifier(value: str, label: str) -> str: 

84 if not value or not IDENTIFIER_RE.match(value): 

85 raise RuntimeError(f'Invalid {label}') 

86 return value 

87 

88 

89async def _retrieve_qdrant(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]: 

90 try: 

91 from qdrant_client import QdrantClient 

92 except ImportError as exc: 

93 raise RuntimeError('qdrant-client is not installed') from exc 

94 

95 if not embedding_function: 

96 raise RuntimeError('Embedding function is not configured') 

97 

98 config = connection.get('config') or {} 

99 external = (knowledge.meta or {}).get('external', {}) 

100 source = external.get('source') or {} 

101 collection_name = source.get('name') 

102 if not collection_name: 

103 raise RuntimeError('External source collection is not configured') 

104 source_config = _source_config(knowledge) 

105 vector_field = source_config.get('vector_field') or None 

106 

107 vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX) 

108 

109 def _search(): 

110 client = QdrantClient( 

111 url=connection.get('endpoint'), 

112 api_key=(auth_config or {}).get('api_key'), 

113 timeout=config.get('timeout') or 30, 

114 ) 

115 return client.query_points( 

116 collection_name=collection_name, 

117 query=vector, 

118 using=vector_field, 

119 limit=count, 

120 ) 

121 

122 response = await asyncio.to_thread(_search) 

123 mapping = { 

124 'content_field': source_config.get('content_field') or 'payload.text', 

125 'metadata_field': source_config.get('metadata_field') or 'payload.metadata', 

126 'document_id_field': source_config.get('document_id_field') or 'id', 

127 'score_field': 'score', 

128 } 

129 

130 normalized = [] 

131 for point in response.points: 

132 normalized.append(_normalize_result(point.model_dump(), mapping, knowledge, distance=point.score)) 

133 return normalized 

134 

135 

136async def _retrieve_milvus(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]: 

137 try: 

138 from pymilvus import MilvusClient 

139 except ImportError as exc: 

140 raise RuntimeError('pymilvus is not installed') from exc 

141 

142 if not embedding_function: 

143 raise RuntimeError('Embedding function is not configured') 

144 

145 config = connection.get('config') or {} 

146 external = (knowledge.meta or {}).get('external', {}) 

147 source = external.get('source') or {} 

148 collection_name = source.get('name') 

149 if not collection_name: 

150 raise RuntimeError('Milvus collection is not configured') 

151 source_config = _source_config(knowledge) 

152 vector_field = source_config.get('vector_field') or 'vector' 

153 content_field = source_config.get('content_field') or 'data.text' 

154 metadata_field = source_config.get('metadata_field') or 'metadata' 

155 

156 vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX) 

157 

158 def _search(): 

159 client_kwargs = { 

160 'uri': connection.get('endpoint'), 

161 } 

162 token = (auth_config or {}).get('api_key') or (auth_config or {}).get('token') 

163 if token: 

164 client_kwargs['token'] = token 

165 if config.get('db_name'): 

166 client_kwargs['db_name'] = config.get('db_name') 

167 

168 client = MilvusClient(**client_kwargs) 

169 output_fields = { 

170 field 

171 for field in ( 

172 _root_field(content_field), 

173 _root_field(metadata_field), 

174 _root_field(source_config.get('document_id_field')), 

175 ) 

176 if field and field != vector_field 

177 } 

178 kwargs = { 

179 'collection_name': collection_name, 

180 'data': [vector], 

181 'anns_field': vector_field, 

182 'limit': count, 

183 'output_fields': list(output_fields), 

184 } 

185 return client.search(**kwargs) 

186 

187 response = await asyncio.to_thread(_search) 

188 mapping = { 

189 'content_field': content_field, 

190 'metadata_field': metadata_field, 

191 'document_id_field': source_config.get('document_id_field') or 'id', 

192 'score_field': 'distance', 

193 } 

194 

195 normalized = [] 

196 for hit in response[0] if response else []: 

197 item = dict(hit) 

198 entity = item.get('entity') or {} 

199 result = { 

200 **entity, 

201 'id': item.get('id') or entity.get('id'), 

202 'distance': item.get('distance'), 

203 } 

204 normalized.append(_normalize_result(result, mapping, knowledge, distance=item.get('distance'))) 

205 return normalized 

206 

207 

208async def _retrieve_pgvector(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]: 

209 try: 

210 import psycopg 

211 from pgvector.psycopg import register_vector 

212 from psycopg.rows import dict_row 

213 except ImportError as exc: 

214 raise RuntimeError('psycopg and pgvector are required for pgvector retrieval') from exc 

215 

216 if not embedding_function: 

217 raise RuntimeError('Embedding function is not configured') 

218 

219 config = connection.get('config') or {} 

220 external = (knowledge.meta or {}).get('external', {}) 

221 source = external.get('source') or {} 

222 collection_name = source.get('name') 

223 if not collection_name: 

224 raise RuntimeError('pgvector collection is not configured') 

225 source_config = _source_config(knowledge) 

226 table_name = source_config.get('table_name') or 'document_chunk' 

227 collection_field = source_config.get('collection_field') or 'collection_name' 

228 content_field = source_config.get('content_field') or 'text' 

229 vector_field = source_config.get('vector_field') or 'vector' 

230 metadata_field = source_config.get('metadata_field') or 'vmetadata' 

231 document_id_field = source_config.get('document_id_field') or 'id' 

232 

233 vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX) 

234 

235 def _search(): 

236 from psycopg import sql 

237 

238 table_identifier = sql.SQL('.').join( 

239 sql.Identifier(_safe_identifier(part, 'table name')) for part in table_name.split('.') 

240 ) 

241 collection_identifier = sql.Identifier(_safe_identifier(collection_field, 'collection field')) 

242 content_identifier = sql.Identifier(_safe_identifier(content_field, 'content field')) 

243 vector_identifier = sql.Identifier(_safe_identifier(vector_field, 'vector field')) 

244 document_id_identifier = sql.Identifier(_safe_identifier(document_id_field, 'document id field')) 

245 metadata_sql = ( 

246 sql.Identifier(_safe_identifier(metadata_field, 'metadata field')) 

247 if metadata_field 

248 else sql.SQL("'{}'::jsonb") 

249 ) 

250 

251 with psycopg.connect( 

252 connection.get('endpoint'), 

253 row_factory=dict_row, 

254 connect_timeout=config.get('timeout') or 30, 

255 ) as conn: 

256 register_vector(conn) 

257 with conn.cursor() as cur: 

258 cur.execute( 

259 sql.SQL( 

260 """ 

261 SELECT {document_id} AS id, 

262 {content} AS content, 

263 {metadata} AS metadata, 

264 {vector_column} <=> %s AS distance 

265 FROM {table_name} 

266 WHERE {collection} = %s 

267 ORDER BY distance ASC 

268 LIMIT %s 

269 """ 

270 ).format( 

271 document_id=document_id_identifier, 

272 content=content_identifier, 

273 metadata=metadata_sql, 

274 vector_column=vector_identifier, 

275 table_name=table_identifier, 

276 collection=collection_identifier, 

277 ), 

278 (vector, collection_name, count), 

279 ) 

280 return cur.fetchall() 

281 

282 rows = await asyncio.to_thread(_search) 

283 mapping = { 

284 'content_field': 'content', 

285 'metadata_field': 'metadata', 

286 'document_id_field': 'id', 

287 'score_field': 'distance', 

288 } 

289 return [_normalize_result(row, mapping, knowledge, distance=row.get('distance')) for row in rows] 

290 

291 

292async def retrieve_external_knowledge( 

293 request, 

294 knowledge: KnowledgeModel, 

295 queries: list[str], 

296 count: int, 

297 user=None, 

298) -> dict: 

299 external = (knowledge.meta or {}).get('external', {}) 

300 connection_id = external.get('connection_id') 

301 if not connection_id: 

302 raise RuntimeError('External knowledge connection is not configured') 

303 

304 connection = await _get_external_connection(connection_id) 

305 if not connection: 

306 raise RuntimeError('External knowledge connection not found') 

307 

308 return await retrieve_external_knowledge_for_connection(request, knowledge, connection, queries, count, user=user) 

309 

310 

311async def retrieve_external_knowledge_for_connection( 

312 request, 

313 knowledge: KnowledgeModel, 

314 connection: dict, 

315 queries: list[str], 

316 count: int, 

317 user=None, 

318) -> dict: 

319 auth_config = connection.get('auth_config') or {} 

320 if not connection.get('enabled', True): 

321 raise RuntimeError('External knowledge connection is disabled') 

322 

323 started_at = time.monotonic() 

324 chunks = [] 

325 provider = (connection.get('provider') or '').lower() 

326 

327 for query in queries: 

328 if provider == 'qdrant': 

329 chunks.extend( 

330 await _retrieve_qdrant( 

331 connection, 

332 auth_config, 

333 knowledge, 

334 query, 

335 count, 

336 getattr(request.app.state, 'EMBEDDING_FUNCTION', None), 

337 ) 

338 ) 

339 elif provider == 'milvus': 

340 chunks.extend( 

341 await _retrieve_milvus( 

342 connection, 

343 auth_config, 

344 knowledge, 

345 query, 

346 count, 

347 getattr(request.app.state, 'EMBEDDING_FUNCTION', None), 

348 ) 

349 ) 

350 elif provider == 'pgvector': 

351 chunks.extend( 

352 await _retrieve_pgvector( 

353 connection, 

354 auth_config, 

355 knowledge, 

356 query, 

357 count, 

358 getattr(request.app.state, 'EMBEDDING_FUNCTION', None), 

359 ) 

360 ) 

361 else: 

362 raise RuntimeError(f'Unsupported external knowledge provider: {connection.get("provider")}') 

363 

364 chunks = chunks[:count] 

365 log.info( 

366 'external_knowledge_retrieval knowledge_id=%s connection_id=%s provider=%s user_id=%s latency_ms=%s result_count=%s', 

367 knowledge.id, 

368 connection.get('id'), 

369 connection.get('provider'), 

370 getattr(user, 'id', None), 

371 round((time.monotonic() - started_at) * 1000), 

372 len(chunks), 

373 ) 

374 

375 return { 

376 'documents': [[chunk['content'] for chunk in chunks]], 

377 'metadatas': [[chunk['metadata'] for chunk in chunks]], 

378 'distances': [[chunk['distance'] for chunk in chunks]], 

379 }