Coverage for open_webui/retrieval/utils.py: 16%

810 statements  

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

1from __future__ import annotations 

2 

3import asyncio 

4import hashlib 

5import logging 

6import os 

7import re 

8import time 

9from typing import Awaitable, Optional, Union 

10from urllib.parse import quote 

11 

12import aiohttp 

13import numpy as np 

14import requests 

15from fastapi import HTTPException 

16from langchain_classic.retrievers import ( 

17 ContextualCompressionRetriever, 

18 EnsembleRetriever, 

19) 

20from langchain_core.documents import Document 

21from open_webui.config import ( 

22 RAG_EMBEDDING_CONTENT_PREFIX, 

23 RAG_EMBEDDING_PREFIX_FIELD_NAME, 

24 RAG_EMBEDDING_QUERY_PREFIX, 

25 VECTOR_DB, 

26) 

27from open_webui.constants import ERROR_MESSAGES 

28from open_webui.env import ( 

29 AIOHTTP_CLIENT_ALLOW_REDIRECTS, 

30 AIOHTTP_CLIENT_SESSION_SSL, 

31 AIOHTTP_CLIENT_TIMEOUT, 

32 BYPASS_RETRIEVAL_ACCESS_CONTROL, 

33 ENABLE_FORWARD_USER_INFO_HEADERS, 

34 ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS, 

35 MPS_INFERENCE_LOCK, 

36 OFFLINE_MODE, 

37 RAG_SOURCE_METADATA_KEYS, 

38 USE_SLIM, 

39) 

40from open_webui.models.access_grants import AccessGrants 

41from open_webui.models.chats import Chats 

42from open_webui.models.files import Files 

43from open_webui.models.folders import Folders 

44from open_webui.models.knowledge import Knowledges 

45from open_webui.models.notes import Notes 

46from open_webui.models.config import Config 

47from open_webui.models.users import UserModel 

48from open_webui.retrieval.loaders.youtube import YoutubeLoader 

49from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT 

50from open_webui.retrieval.external import retrieve_external_knowledge 

51from open_webui.retrieval.vector.factory import get_vector_db_client 

52from open_webui.retrieval.vector.main import GetResult, SearchResult 

53from open_webui.retrieval.web.utils import get_web_loader 

54from open_webui.utils.access_control.files import get_owner_accessible_folder_files, has_access_to_file 

55from open_webui.utils.access_control.folders import has_folder_access 

56from open_webui.utils.headers import get_json_bearer_headers, include_user_info_headers 

57from open_webui.utils.misc import get_content_from_message, get_message_list 

58 

59log = logging.getLogger(__name__) 

60 

61 

62from typing import Any 

63 

64from langchain_core.callbacks import CallbackManagerForRetrieverRun 

65from langchain_core.retrievers import BaseRetriever 

66 

67 

68class BM25Retriever(BaseRetriever): 

69 docs: list[Document] 

70 vectorizer: Any 

71 k: int 

72 

73 def _get_relevant_documents(self, query: str, *, run_manager: CallbackManagerForRetrieverRun) -> list[Document]: 

74 return self.vectorizer.get_top_n(query.split(), self.docs, n=self.k) 

75 

76 

77def is_youtube_url(url: str) -> bool: 

78 youtube_regex = r'^(https?://)?(www\.)?(youtube\.com|youtu\.be)/.+$' 

79 return re.match(youtube_regex, url) is not None 

80 

81 

82LOADER_CONFIG_KEYS = { 

83 'file_max_size': 'rag.file.max_size', 

84 'youtube_language': 'rag.youtube_loader_language', 

85 'youtube_proxy_url': 'rag.youtube_loader_proxy_url', 

86 'web_loader_ssl_verification': 'web.loader.ssl_verification', 

87 'web_loader_concurrent_requests': 'web.loader.concurrent_requests', 

88 'web_search_trust_env': 'web.search.trust_env', 

89 'web_loader_engine': 'web.loader.engine', 

90 'web_loader_timeout': 'web.loader.timeout', 

91 'playwright_ws_url': 'web.loader.playwright_ws_url', 

92 'playwright_timeout': 'web.loader.playwright_timeout', 

93 'firecrawl_api_key': 'web.loader.firecrawl_api_key', 

94 'firecrawl_api_url': 'web.loader.firecrawl_api_url', 

95 'firecrawl_timeout': 'web.loader.firecrawl_timeout', 

96 'tavily_api_key': 'web.search.tavily_api_key', 

97 'tavily_extract_depth': 'web.search.tavily_extract_depth', 

98 'microsoft_web_iq_api_base_url': 'web.search.microsoft_web_iq_api_base_url', 

99 'microsoft_web_iq_api_key': 'web.search.microsoft_web_iq_api_key', 

100 'microsoft_web_iq_language': 'web.search.microsoft_web_iq_language', 

101 'external_web_loader_url': 'web.loader.external_web_loader_url', 

102 'external_web_loader_api_key': 'web.loader.external_web_loader_api_key', 

103 'CONTENT_EXTRACTION_ENGINE': 'rag.content_extraction_engine', 

104 'DATALAB_MARKER_API_KEY': 'rag.datalab_marker_api_key', 

105 'DATALAB_MARKER_API_BASE_URL': 'rag.datalab_marker_api_base_url', 

106 'DATALAB_MARKER_ADDITIONAL_CONFIG': 'rag.datalab_marker_additional_config', 

107 'DATALAB_MARKER_SKIP_CACHE': 'rag.datalab_marker_skip_cache', 

108 'DATALAB_MARKER_FORCE_OCR': 'rag.datalab_marker_force_ocr', 

109 'DATALAB_MARKER_PAGINATE': 'rag.datalab_marker_paginate', 

110 'DATALAB_MARKER_STRIP_EXISTING_OCR': 'rag.datalab_marker_strip_existing_ocr', 

111 'DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION': 'rag.datalab_marker_disable_image_extraction', 

112 'DATALAB_MARKER_FORMAT_LINES': 'rag.datalab_marker_format_lines', 

113 'DATALAB_MARKER_USE_LLM': 'rag.datalab_marker_use_llm', 

114 'DATALAB_MARKER_OUTPUT_FORMAT': 'rag.datalab_marker_output_format', 

115 'EXTERNAL_DOCUMENT_LOADER_URL': 'rag.external_document_loader_url', 

116 'EXTERNAL_DOCUMENT_LOADER_API_KEY': 'rag.external_document_loader_api_key', 

117 'EXTERNAL_DOCUMENT_LOADER_HEADERS': 'rag.external_document_loader_headers', 

118 'TIKA_SERVER_URL': 'rag.tika_server_url', 

119 'TIKA_SERVER_VERSION': 'rag.tika_server_version', 

120 'DOCLING_SERVER_URL': 'rag.docling_server_url', 

121 'DOCLING_API_KEY': 'rag.docling_api_key', 

122 'DOCLING_PARAMS': 'rag.docling_params', 

123 'PDF_EXTRACT_IMAGES': 'rag.pdf_extract_images', 

124 'PDF_LOADER_MODE': 'rag.pdf_loader_mode', 

125 'DOCUMENT_INTELLIGENCE_ENDPOINT': 'rag.document_intelligence_endpoint', 

126 'DOCUMENT_INTELLIGENCE_KEY': 'rag.document_intelligence_key', 

127 'DOCUMENT_INTELLIGENCE_MODEL': 'rag.document_intelligence_model', 

128 'MISTRAL_OCR_API_BASE_URL': 'rag.mistral_ocr_api_base_url', 

129 'MISTRAL_OCR_API_KEY': 'rag.mistral_ocr_api_key', 

130 'MISTRAL_OCR_USE_BASE64': 'rag.mistral_ocr_use_base64', 

131 'PADDLEOCR_VL_BASE_URL': 'rag.paddleocr_vl_base_url', 

132 'PADDLEOCR_VL_TOKEN': 'rag.paddleocr_vl_token', 

133 'MINERU_API_MODE': 'rag.mineru_api_mode', 

134 'MINERU_API_URL': 'rag.mineru_api_url', 

135 'MINERU_API_KEY': 'rag.mineru_api_key', 

136 'MINERU_API_TIMEOUT': 'rag.mineru_api_timeout', 

137 'MINERU_PARAMS': 'rag.mineru_params', 

138 'MINERU_FILE_EXTENSIONS': 'rag.mineru_file_extensions', 

139} 

140 

141 

142async def get_loader_config(): 

143 values = await Config.get_many(*LOADER_CONFIG_KEYS.values()) 

144 return {name: values.get(key) for name, key in LOADER_CONFIG_KEYS.items()} 

145 

146 

147def get_loader(request, url: str, config: dict): 

148 if is_youtube_url(url): 

149 return YoutubeLoader( 

150 url, 

151 language=config.get('youtube_language'), 

152 proxy_url=config.get('youtube_proxy_url'), 

153 ) 

154 return get_web_loader( 

155 url, 

156 verify_ssl=config.get('web_loader_ssl_verification'), 

157 requests_per_second=config.get('web_loader_concurrent_requests'), 

158 trust_env=config.get('web_search_trust_env'), 

159 loader_config=config, 

160 ) 

161 

162 

163def build_loader_from_config(request, config: dict): 

164 """Build a Loader instance with the admin's configured extraction engine settings.""" 

165 from open_webui.retrieval.loaders.main import Loader 

166 

167 loader_config = {key: config.get(key) for key in LOADER_CONFIG_KEYS if key.isupper()} 

168 loader_config['FILE_MAX_SIZE'] = config.get('file_max_size') 

169 return Loader( 

170 engine=loader_config['CONTENT_EXTRACTION_ENGINE'], 

171 **{key: value for key, value in loader_config.items() if key != 'CONTENT_EXTRACTION_ENGINE'}, 

172 ) 

173 

174 

175def _extract_text_from_binary_response( 

176 request, response: requests.Response, url: str, loader_config: dict 

177) -> tuple[str, list]: 

178 """Download response body to a temp file and extract text using the Loader pipeline.""" 

179 import mimetypes 

180 import tempfile 

181 import urllib.parse 

182 

183 content_type = response.headers.get('Content-Type', '').split(';')[0].strip() 

184 

185 # Derive filename from URL path, falling back to Content-Disposition or mime guess 

186 url_path = urllib.parse.urlparse(url).path 

187 filename = os.path.basename(url_path) if url_path else '' 

188 

189 if not filename or '.' not in filename: 

190 # Try Content-Disposition header 

191 cd = response.headers.get('Content-Disposition', '') 

192 if 'filename=' in cd: 

193 filename = cd.split('filename=')[-1].strip('"\'') 

194 

195 if not filename or '.' not in filename: 

196 ext = mimetypes.guess_extension(content_type) or '' 

197 filename = f'download{ext}' 

198 

199 suffix = '.' + filename.split('.')[-1].lower() if '.' in filename else '' 

200 

201 max_size = loader_config.get('file_max_size') 

202 max_bytes = int(max_size) * 1024 * 1024 if max_size else 0 

203 

204 tmp_fd, tmp_path = tempfile.mkstemp(suffix=suffix) 

205 try: 

206 downloaded = 0 

207 # Stream to disk; response.content buffers the whole body in memory first. 

208 with os.fdopen(tmp_fd, 'wb') as tmp: 

209 for chunk in response.iter_content(64 * 1024): 

210 downloaded += len(chunk) 

211 if max_bytes and downloaded > max_bytes: 

212 raise ValueError(ERROR_MESSAGES.FILE_TOO_LARGE(size=f'{max_size} MB')) 

213 tmp.write(chunk) 

214 

215 loader = build_loader_from_config(request, loader_config) 

216 docs = loader.load(filename, content_type, tmp_path) 

217 for doc in docs: 

218 doc.metadata['source'] = url 

219 content = ' '.join([doc.page_content for doc in docs]) 

220 return content, docs 

221 finally: 

222 os.remove(tmp_path) 

223 

224 

225TEXT_APPLICATION_CONTENT_TYPES = { 

226 'application/javascript', 

227 'application/json', 

228 'application/xml', 

229 'application/x-javascript', 

230} 

231 

232 

233def _is_text_content_type(content_type: str) -> bool: 

234 """Return True if the content type should be handled by the web loader.""" 

235 ct = content_type.split(';')[0].strip().lower() 

236 if not ct: 

237 return True 

238 if ct.startswith('text/'): 

239 return True 

240 if ct in TEXT_APPLICATION_CONTENT_TYPES: 

241 return True 

242 return ct.endswith(('+xml', '+json')) 

243 

244 

245async def get_content_from_url(request, url: str) -> str: 

246 loader_config = await get_loader_config() 

247 

248 # The rest of this function performs synchronous, blocking work: an SSRF-guarded 

249 # `requests` probe and a synchronous document loader (`loader.load()`). Run it in a 

250 # worker thread so the event loop stays free while waiting on network/parsing. 

251 return await asyncio.to_thread(_get_content_from_url_sync, request, url, loader_config) 

252 

253 

254def _get_content_from_url_sync(request, url: str, loader_config): 

255 from open_webui.retrieval.web.utils import validate_url, get_ssrf_safe_requests_session 

256 

257 # Validate URL before making any request (blocks private IPs, non-HTTP, filter list) 

258 validate_url(url) 

259 

260 # YouTube URLs (including youtu.be short links) should go straight to 

261 # YoutubeLoader, which uses youtube-transcript-api and never needs the 

262 # HTTP response body. Probing the URL first is harmful for short URLs: 

263 # youtu.be returns a 303 redirect with Content-Type: application/binary 

264 # when allow_redirects=False, causing the binary-content path to run 

265 # and produce empty docs → HTTP 400. 

266 if is_youtube_url(url): 

267 loader = get_loader(request, url, loader_config) 

268 docs = loader.load() 

269 content = ' '.join([doc.page_content for doc in docs]) 

270 return content, docs 

271 

272 # Streamed GET to check Content-Type without downloading the body. 

273 # allow_redirects=False prevents redirect-based SSRF: validate_url() above is 

274 # called on the originally-submitted URL only; following 3xx redirects without 

275 # re-validation would let an attacker reach private IPs (RFC1918, loopback, 

276 # cloud-metadata 169.254.169.254) via a public host that redirects internally. 

277 try: 

278 # Probe through the connect-time SSRF guard; bare requests.get re-resolves (DNS-rebinding gap). 

279 session = get_ssrf_safe_requests_session() 

280 response = session.get(url, stream=True, timeout=30, allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS) 

281 response.raise_for_status() 

282 content_type = response.headers.get('Content-Type', '') 

283 except Exception: 

284 content_type = '' 

285 response = None 

286 

287 # Text / HTML / unknown — use the configured web loader 

288 if response is None or _is_text_content_type(content_type): 

289 if response is not None: 

290 response.close() 

291 loader = get_loader(request, url, loader_config) 

292 docs = loader.load() 

293 content = ' '.join([doc.page_content for doc in docs]) 

294 return content, docs 

295 

296 # Binary content (PDF, DOCX, XLSX, PPTX, etc.) — download and extract 

297 try: 

298 return _extract_text_from_binary_response(request, response, url, loader_config) 

299 finally: 

300 response.close() 

301 

302 

303CHUNK_HASH_KEY = '_chunk_hash' 

304 

305 

306def _content_hash(text: str) -> str: 

307 """SHA-256 hash of text, used as a stable chunk identifier for RRF dedup.""" 

308 return hashlib.sha256(text.encode()).hexdigest() 

309 

310 

311class VectorSearchRetriever(BaseRetriever): 

312 collection_name: Any 

313 embedding_function: Any 

314 top_k: int 

315 

316 def _get_relevant_documents(self, query: str, *, run_manager: CallbackManagerForRetrieverRun) -> list[Document]: 

317 """Get documents relevant to a query. 

318 

319 Args: 

320 query: String to find relevant documents for. 

321 run_manager: The callback handler to use. 

322 

323 Returns: 

324 List of relevant documents. 

325 """ 

326 return [] 

327 

328 async def _aget_relevant_documents( 

329 self, 

330 query: str, 

331 *, 

332 run_manager: CallbackManagerForRetrieverRun, 

333 ) -> list[Document]: 

334 embedding = await self.embedding_function(query, RAG_EMBEDDING_QUERY_PREFIX) 

335 result = await ASYNC_VECTOR_DB_CLIENT.search( 

336 collection_name=self.collection_name, 

337 vectors=[embedding], 

338 limit=self.top_k, 

339 ) 

340 

341 return _search_result_to_documents(result) 

342 

343 

344def query_doc(collection_name: str, query_embedding: list[float], k: int, user: UserModel = None): 

345 log.debug('query_doc:doc %s', collection_name) 

346 result = get_vector_db_client().search( 

347 collection_name=collection_name, 

348 vectors=[query_embedding], 

349 limit=k, 

350 ) 

351 

352 if result: 

353 log.info('query_doc:result %s %s', result.ids, result.metadatas) 

354 

355 return result 

356 

357 

358def get_doc(collection_name: str, user: UserModel = None): 

359 try: 

360 log.debug('get_doc:doc %s', collection_name) 

361 result = get_vector_db_client().get(collection_name=collection_name) 

362 

363 if result: 

364 log.info('query_doc:result %s %s', result.ids, result.metadatas) 

365 

366 return result 

367 except Exception as e: 

368 log.exception(f'Error getting doc {collection_name}: {e}') 

369 raise e 

370 

371 

372def get_enriched_texts(collection_result: GetResult) -> list[str]: 

373 enriched_texts = [] 

374 for idx, text in enumerate(collection_result.documents[0]): 

375 metadata = collection_result.metadatas[0][idx] 

376 metadata_parts = [text] 

377 

378 # Add filename (repeat twice for extra weight in BM25 scoring) 

379 if metadata.get('name'): 

380 filename = metadata['name'] 

381 filename_tokens = filename.replace('_', ' ').replace('-', ' ').replace('.', ' ') 

382 metadata_parts.append(f'Filename: {filename} {filename_tokens} {filename_tokens}') 

383 

384 # Add title if available 

385 if metadata.get('title'): 

386 metadata_parts.append(f'Title: {metadata["title"]}') 

387 

388 # Add document section headings if available (from markdown splitter) 

389 if metadata.get('headings') and isinstance(metadata['headings'], list): 

390 headings = ' > '.join(str(h) for h in metadata['headings']) 

391 metadata_parts.append(f'Section: {headings}') 

392 

393 # Add source URL/path if available 

394 if metadata.get('source'): 

395 metadata_parts.append(f'Source: {metadata["source"]}') 

396 

397 # Add snippet for web search results 

398 if metadata.get('snippet'): 

399 metadata_parts.append(f'Snippet: {metadata["snippet"]}') 

400 

401 enriched_texts.append(' '.join(metadata_parts)) 

402 

403 return enriched_texts 

404 

405 

406def _search_result_to_documents(result: SearchResult | None) -> list[Document]: 

407 ids = result.ids[0] if result and result.ids else [] 

408 metadatas = result.metadatas[0] if result and result.metadatas else [] 

409 documents = result.documents[0] if result and result.documents else [] 

410 distances = result.distances[0] if result and result.distances else [] 

411 

412 docs = [] 

413 for idx in range(len(ids)): 

414 document = documents[idx] 

415 metadata = dict(metadatas[idx] or {}) 

416 metadata[CHUNK_HASH_KEY] = _content_hash(document) 

417 if idx < len(distances): 

418 metadata.setdefault('score', distances[idx]) 

419 docs.append(Document(metadata=metadata, page_content=document)) 

420 return docs 

421 

422 

423def _supports_native_hybrid_search() -> bool: 

424 supports_hybrid_search = getattr(ASYNC_VECTOR_DB_CLIENT, 'supports_hybrid_search', None) 

425 if supports_hybrid_search is not None: 

426 return bool(supports_hybrid_search) 

427 return callable(getattr(ASYNC_VECTOR_DB_CLIENT, 'hybrid_search', None)) 

428 

429 

430async def query_doc_with_native_hybrid_search( 

431 collection_name: str, 

432 query: str, 

433 embedding_function, 

434 k: int, 

435 reranking_function, 

436 k_reranker: int, 

437 r: float, 

438 hybrid_bm25_weight: float, 

439) -> Optional[dict]: 

440 try: 

441 if not _supports_native_hybrid_search(): 

442 return None 

443 

444 query_vectors = [] 

445 if hybrid_bm25_weight < 1: 

446 query_vectors = [await embedding_function(query, RAG_EMBEDDING_QUERY_PREFIX)] 

447 

448 result = await ASYNC_VECTOR_DB_CLIENT.hybrid_search( 

449 collection_name=collection_name, 

450 query=query, 

451 vectors=query_vectors, 

452 limit=k, 

453 hybrid_bm25_weight=hybrid_bm25_weight, 

454 ) 

455 if result is None: 

456 return None 

457 

458 documents = _search_result_to_documents(result) 

459 if not documents: 

460 return {'distances': [[]], 'documents': [[]], 'metadatas': [[]]} 

461 

462 compressor = RerankCompressor( 

463 embedding_function=embedding_function, 

464 top_n=k_reranker, 

465 reranking_function=reranking_function, 

466 r_score=r, 

467 ) 

468 compressed = await compressor.acompress_documents(documents, query) 

469 

470 distances = [d.metadata.get('score') for d in compressed] 

471 documents = [d.page_content for d in compressed] 

472 metadatas = [d.metadata for d in compressed] 

473 

474 if k < k_reranker: 

475 sorted_items = sorted(zip(distances, documents, metadatas), key=lambda x: x[0], reverse=True) 

476 sorted_items = sorted_items[:k] 

477 

478 if sorted_items: 

479 distances, documents, metadatas = map(list, zip(*sorted_items)) 

480 else: 

481 distances, documents, metadatas = [], [], [] 

482 

483 return { 

484 'distances': [distances], 

485 'documents': [documents], 

486 'metadatas': [metadatas], 

487 } 

488 except Exception as e: 

489 log.debug('Native hybrid search failed for %s, falling back to legacy hybrid search: %s', collection_name, e) 

490 return None 

491 

492 

493async def query_doc_with_hybrid_search( 

494 collection_name: str, 

495 collection_result: Optional[GetResult], 

496 query: str, 

497 embedding_function, 

498 k: int, 

499 reranking_function, 

500 k_reranker: int, 

501 r: float, 

502 hybrid_bm25_weight: float, 

503 enable_enriched_texts: bool = False, 

504 native_hybrid_search: bool = True, 

505) -> dict: 

506 if native_hybrid_search and not enable_enriched_texts: 

507 native_result = await query_doc_with_native_hybrid_search( 

508 collection_name=collection_name, 

509 query=query, 

510 embedding_function=embedding_function, 

511 k=k, 

512 reranking_function=reranking_function, 

513 k_reranker=k_reranker, 

514 r=r, 

515 hybrid_bm25_weight=hybrid_bm25_weight, 

516 ) 

517 if native_result is not None: 

518 return native_result 

519 

520 if collection_result is None: 

521 collection_result = await ASYNC_VECTOR_DB_CLIENT.get(collection_name=collection_name) 

522 

523 # First check if collection_result has the required attributes 

524 if ( 

525 not collection_result 

526 or not hasattr(collection_result, 'documents') 

527 or not hasattr(collection_result, 'metadatas') 

528 ): 

529 log.warning(f'query_doc_with_hybrid_search:no_docs {collection_name}') 

530 return {'documents': [], 'metadatas': [], 'distances': []} 

531 

532 # Now safely check the documents content after confirming attributes exist 

533 if not collection_result.documents or len(collection_result.documents) == 0 or not collection_result.documents[0]: 

534 log.warning(f'query_doc_with_hybrid_search:no_docs {collection_name}') 

535 return {'documents': [], 'metadatas': [], 'distances': []} 

536 

537 log.debug('query_doc_with_hybrid_search:doc %s', collection_name) 

538 

539 original_texts = collection_result.documents[0] 

540 bm25_metadatas = [ 

541 {**meta, CHUNK_HASH_KEY: _content_hash(original_texts[idx])} 

542 for idx, meta in enumerate(collection_result.metadatas[0]) 

543 ] 

544 

545 bm25_texts = get_enriched_texts(collection_result) if enable_enriched_texts else original_texts 

546 

547 from rank_bm25 import BM25Okapi 

548 

549 bm25_retriever = BM25Retriever( 

550 docs=[Document(page_content=text, metadata=meta) for text, meta in zip(bm25_texts, bm25_metadatas)], 

551 vectorizer=BM25Okapi([text.split() for text in bm25_texts]), 

552 k=k, 

553 ) 

554 

555 vector_search_retriever = VectorSearchRetriever( 

556 collection_name=collection_name, 

557 embedding_function=embedding_function, 

558 top_k=k, 

559 ) 

560 

561 # Use CHUNK_HASH_KEY for dedup so enriched BM25 texts don't defeat RRF 

562 if hybrid_bm25_weight <= 0: 

563 ensemble_retriever = EnsembleRetriever( 

564 retrievers=[vector_search_retriever], 

565 weights=[1.0], 

566 id_key=CHUNK_HASH_KEY, 

567 ) 

568 elif hybrid_bm25_weight >= 1: 

569 ensemble_retriever = EnsembleRetriever( 

570 retrievers=[bm25_retriever], 

571 weights=[1.0], 

572 id_key=CHUNK_HASH_KEY, 

573 ) 

574 else: 

575 ensemble_retriever = EnsembleRetriever( 

576 retrievers=[bm25_retriever, vector_search_retriever], 

577 weights=[hybrid_bm25_weight, 1.0 - hybrid_bm25_weight], 

578 id_key=CHUNK_HASH_KEY, 

579 ) 

580 

581 compressor = RerankCompressor( 

582 embedding_function=embedding_function, 

583 top_n=k_reranker, 

584 reranking_function=reranking_function, 

585 r_score=r, 

586 ) 

587 

588 compression_retriever = ContextualCompressionRetriever( 

589 base_compressor=compressor, base_retriever=ensemble_retriever 

590 ) 

591 

592 result = await compression_retriever.ainvoke(query) 

593 

594 distances = [d.metadata.get('score') for d in result] 

595 documents = [d.page_content for d in result] 

596 metadatas = [d.metadata for d in result] 

597 

598 # retrieve only min(k, k_reranker) items, sort and cut by distance if k < k_reranker 

599 if k < k_reranker: 

600 sorted_items = sorted(zip(distances, documents, metadatas), key=lambda x: x[0], reverse=True) 

601 sorted_items = sorted_items[:k] 

602 

603 if sorted_items: 

604 distances, documents, metadatas = map(list, zip(*sorted_items)) 

605 else: 

606 distances, documents, metadatas = [], [], [] 

607 

608 result = { 

609 'distances': [distances], 

610 'documents': [documents], 

611 'metadatas': [metadatas], 

612 } 

613 

614 log.info('query_doc_with_hybrid_search:result %s %s', result['metadatas'], result['distances']) 

615 return result 

616 

617 

618def merge_get_results(get_results: list[dict]) -> dict: 

619 # Initialize lists to store combined data 

620 combined_documents = [] 

621 combined_metadatas = [] 

622 combined_ids = [] 

623 

624 for data in get_results: 

625 combined_documents.extend(data['documents'][0]) 

626 combined_metadatas.extend(data['metadatas'][0]) 

627 combined_ids.extend(data['ids'][0]) 

628 

629 # Create the output dictionary 

630 result = { 

631 'documents': [combined_documents], 

632 'metadatas': [combined_metadatas], 

633 'ids': [combined_ids], 

634 } 

635 

636 return result 

637 

638 

639def merge_and_sort_query_results(query_results: list[dict], k: int) -> dict: 

640 # Initialize lists to store combined data 

641 combined = dict() # To store documents with unique document hashes 

642 

643 for data in query_results: 

644 if ( 

645 len(data.get('distances', [])) == 0 

646 or len(data.get('documents', [])) == 0 

647 or len(data.get('metadatas', [])) == 0 

648 ): 

649 continue 

650 

651 distances = data['distances'][0] 

652 documents = data['documents'][0] 

653 metadatas = data['metadatas'][0] 

654 

655 for distance, document, metadata in zip(distances, documents, metadatas): 

656 if isinstance(document, str): 

657 doc_hash = (metadata or {}).get(CHUNK_HASH_KEY) or _content_hash(document) 

658 

659 if doc_hash not in combined: 

660 combined[doc_hash] = (distance, document, metadata) 

661 continue # if doc is new, no further comparison is needed 

662 

663 # if doc is alredy in, but new distance is better, update 

664 if distance > combined[doc_hash][0]: 

665 combined[doc_hash] = (distance, document, metadata) 

666 

667 combined = list(combined.values()) 

668 # Sort the list based on distances 

669 combined.sort(key=lambda x: x[0], reverse=True) 

670 

671 # Slice to keep only the top k elements 

672 sorted_distances, sorted_documents, sorted_metadatas = zip(*combined[:k]) if combined else ([], [], []) 

673 

674 # Create and return the output dictionary 

675 return { 

676 'distances': [list(sorted_distances)], 

677 'documents': [list(sorted_documents)], 

678 'metadatas': [list(sorted_metadatas)], 

679 } 

680 

681 

682def get_all_items_from_collections(collection_names: list[str]) -> dict: 

683 results = [] 

684 

685 for collection_name in collection_names: 

686 if collection_name: 

687 try: 

688 result = get_doc(collection_name=collection_name) 

689 if result is not None: 

690 results.append(result.model_dump()) 

691 except Exception as e: 

692 log.exception(f'Error when querying the collection: {e}') 

693 else: 

694 pass 

695 

696 return merge_get_results(results) 

697 

698 

699async def query_collection( 

700 request, 

701 collection_names: list[str], 

702 queries: list[str], 

703 embedding_function, 

704 k: int, 

705) -> dict: 

706 config = await Config.get_many( 

707 'rag.enable_hybrid_search', 

708 'rag.top_k_reranker', 

709 'rag.relevance_threshold', 

710 'rag.hybrid_bm25_weight', 

711 'rag.enable_hybrid_search_enriched_texts', 

712 ) 

713 # When request is provided, try hybrid search + reranking if enabled 

714 if request and config.get('rag.enable_hybrid_search'): 714 ↛ 715line 714 didn't jump to line 715 because the condition on line 714 was never true

715 try: 

716 reranking_function = ( 

717 (lambda query, documents: request.app.state.RERANKING_FUNCTION(query, documents)) 

718 if request.app.state.RERANKING_FUNCTION 

719 else None 

720 ) 

721 return await query_collection_with_hybrid_search( 

722 collection_names=collection_names, 

723 queries=queries, 

724 embedding_function=embedding_function, 

725 k=k, 

726 reranking_function=reranking_function, 

727 k_reranker=config.get('rag.top_k_reranker'), 

728 r=config.get('rag.relevance_threshold'), 

729 hybrid_bm25_weight=config.get('rag.hybrid_bm25_weight'), 

730 enable_enriched_texts=config.get('rag.enable_hybrid_search_enriched_texts'), 

731 ) 

732 except Exception as e: 

733 log.debug('Hybrid search failed, falling back to vector search: %s', e) 

734 

735 results = [] 

736 last_error = None 

737 failed_collection_names = set() 

738 

739 def process_query_collection(collection_name, query_embedding): 

740 try: 

741 if collection_name: 

742 result = query_doc( 

743 collection_name=collection_name, 

744 k=k, 

745 query_embedding=query_embedding, 

746 ) 

747 if result is not None: 

748 return result.model_dump(), None, collection_name 

749 return None, None, collection_name 

750 except Exception as e: 

751 return None, e, collection_name 

752 

753 # Sanitize: filter out None/empty queries to prevent embedding crashes 

754 # (e.g. when get_last_user_message returns None) 

755 queries = [q for q in queries if q] 

756 if not queries: 

757 log.warning('query_collection: all queries were None or empty, returning empty results') 

758 return {'distances': [[]], 'documents': [[]], 'metadatas': [[]]} 

759 

760 # Generate all query embeddings (in one call) 

761 query_embeddings = await embedding_function(queries, prefix=RAG_EMBEDDING_QUERY_PREFIX) 

762 log.debug('query_collection: processing %s queries across %s collections', len(queries), len(collection_names)) 

763 

764 task_results = await asyncio.gather( 

765 *[ 

766 asyncio.to_thread(process_query_collection, collection_name, query_embedding) 

767 for query_embedding in query_embeddings 

768 for collection_name in collection_names 

769 ] 

770 ) 

771 

772 for result, err, collection_name in task_results: 

773 if err is not None: 

774 last_error = err 

775 failed_collection_names.add(collection_name) 

776 elif result is not None: 

777 results.append(result) 

778 

779 if failed_collection_names: 

780 log.error( 

781 'query_collection: %s collection(s) had failing queries: %s', 

782 len(failed_collection_names), 

783 ', '.join(sorted(failed_collection_names)), 

784 exc_info=last_error, 

785 ) 

786 

787 return merge_and_sort_query_results(results, k=k) 

788 

789 

790async def query_collection_with_hybrid_search( 

791 collection_names: list[str], 

792 queries: list[str], 

793 embedding_function, 

794 k: int, 

795 reranking_function, 

796 k_reranker: int, 

797 r: float, 

798 hybrid_bm25_weight: float, 

799 enable_enriched_texts: bool = False, 

800) -> dict: 

801 results = [] 

802 last_error = None 

803 failed_collection_names = set() 

804 

805 if not enable_enriched_texts: 

806 

807 async def process_native_query(collection_name, query): 

808 result = await query_doc_with_native_hybrid_search( 

809 collection_name=collection_name, 

810 query=query, 

811 embedding_function=embedding_function, 

812 k=k, 

813 reranking_function=reranking_function, 

814 k_reranker=k_reranker, 

815 r=r, 

816 hybrid_bm25_weight=hybrid_bm25_weight, 

817 ) 

818 return result 

819 

820 native_task_results = await asyncio.gather( 

821 *[process_native_query(collection_name, query) for collection_name in collection_names for query in queries] 

822 ) 

823 if native_task_results and all(result is not None for result in native_task_results): 

824 return merge_and_sort_query_results(native_task_results, k=k) 

825 

826 # Fetch every collection's contents once up front so the 

827 # per-query/per-document loop below can reuse them. Each fetch 

828 # offloads to a worker thread, so run them concurrently with 

829 # `asyncio.gather` instead of awaiting them serially — otherwise 

830 # latency scales linearly with `len(collection_names)`. 

831 log.debug( 

832 'query_collection_with_hybrid_search: prefetching %d collections', 

833 len(collection_names), 

834 ) 

835 

836 async def _fetch_collection(name: str): 

837 try: 

838 return name, await ASYNC_VECTOR_DB_CLIENT.get(collection_name=name) 

839 except Exception as e: 

840 log.exception(f'Failed to fetch collection {name}: {e}') 

841 return name, None 

842 

843 collection_results = dict(await asyncio.gather(*(_fetch_collection(name) for name in collection_names))) 

844 

845 log.info('Starting hybrid search for %s queries in %s collections...', len(queries), len(collection_names)) 

846 

847 async def process_query(collection_name, query): 

848 try: 

849 result = await query_doc_with_hybrid_search( 

850 collection_name=collection_name, 

851 collection_result=collection_results[collection_name], 

852 query=query, 

853 embedding_function=embedding_function, 

854 k=k, 

855 reranking_function=reranking_function, 

856 k_reranker=k_reranker, 

857 r=r, 

858 hybrid_bm25_weight=hybrid_bm25_weight, 

859 enable_enriched_texts=enable_enriched_texts, 

860 native_hybrid_search=False, 

861 ) 

862 return result, None, collection_name 

863 except Exception as e: 

864 return None, e, collection_name 

865 

866 # Prepare tasks for all collections and queries 

867 # Avoid running any tasks for collections that failed to fetch data (have assigned None) 

868 tasks = [ 

869 (collection_name, query) 

870 for collection_name in collection_names 

871 if collection_results[collection_name] is not None 

872 for query in queries 

873 ] 

874 

875 # Run all queries in parallel using asyncio.gather 

876 task_results = await asyncio.gather(*[process_query(collection_name, query) for collection_name, query in tasks]) 

877 

878 for result, err, collection_name in task_results: 

879 if err is not None: 

880 last_error = err 

881 failed_collection_names.add(collection_name) 

882 elif result is not None: 

883 results.append(result) 

884 

885 if failed_collection_names: 

886 log.error( 

887 'query_collection_with_hybrid_search: %s collection(s) had failing queries: %s', 

888 len(failed_collection_names), 

889 ', '.join(sorted(failed_collection_names)), 

890 exc_info=last_error, 

891 ) 

892 

893 if failed_collection_names and not results: 

894 raise Exception('Hybrid search failed for all collections. Using Non-hybrid search as fallback.') 

895 

896 return merge_and_sort_query_results(results, k=k) 

897 

898 

899def generate_openai_batch_embeddings( 

900 model: str, 

901 texts: list[str], 

902 url: str = 'https://api.openai.com/v1', 

903 key: str = '', 

904 prefix: str = None, 

905 user: UserModel = None, 

906) -> list[list[float]]: 

907 log.debug('generate_openai_batch_embeddings:model %s batch size: %s', model, len(texts)) 

908 json_data = {'input': texts, 'model': model} 

909 if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): 

910 json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix 

911 

912 headers = get_json_bearer_headers(key) 

913 if ENABLE_FORWARD_USER_INFO_HEADERS and user: 

914 headers = include_user_info_headers(headers, user) 

915 

916 r = requests.post( 

917 f'{url}/embeddings', 

918 headers=headers, 

919 json=json_data, 

920 ) 

921 r.raise_for_status() 

922 data = r.json() 

923 if 'data' in data: 

924 return [elem['embedding'] for elem in data['data']] 

925 else: 

926 raise ValueError("Unexpected OpenAI embeddings response: missing 'data' key") 

927 

928 

929async def agenerate_openai_batch_embeddings( 

930 model: str, 

931 texts: list[str], 

932 url: str = 'https://api.openai.com/v1', 

933 key: str = '', 

934 prefix: str = None, 

935 user: UserModel = None, 

936) -> list[list[float]]: 

937 log.debug('agenerate_openai_batch_embeddings:model %s batch size: %s', model, len(texts)) 

938 form_data = {'input': texts, 'model': model} 

939 if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): 

940 form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix 

941 

942 headers = get_json_bearer_headers(key) 

943 if ENABLE_FORWARD_USER_INFO_HEADERS and user: 

944 headers = include_user_info_headers(headers, user) 

945 

946 async with aiohttp.ClientSession( 

947 trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) 

948 ) as session: 

949 async with session.post( 

950 f'{url}/embeddings', 

951 headers=headers, 

952 json=form_data, 

953 ssl=AIOHTTP_CLIENT_SESSION_SSL, 

954 ) as r: 

955 r.raise_for_status() 

956 data = await r.json() 

957 if 'data' in data: 

958 return [item['embedding'] for item in data['data']] 

959 else: 

960 raise ValueError("Unexpected OpenAI embeddings response: missing 'data' key") 

961 

962 

963def generate_azure_openai_batch_embeddings( 

964 model: str, 

965 texts: list[str], 

966 url: str, 

967 key: str = '', 

968 version: str = '', 

969 prefix: str = None, 

970 user: UserModel = None, 

971) -> list[list[float]]: 

972 log.debug('generate_azure_openai_batch_embeddings:deployment %s batch size: %s', model, len(texts)) 

973 json_data = {'input': texts} 

974 if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): 

975 json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix 

976 

977 url = f'{url}/openai/deployments/{model}/embeddings?api-version={version}' 

978 

979 for _ in range(5): 

980 headers = { 

981 'Content-Type': 'application/json', 

982 'api-key': key, 

983 } 

984 if ENABLE_FORWARD_USER_INFO_HEADERS and user: 

985 headers = include_user_info_headers(headers, user) 

986 

987 r = requests.post( 

988 url, 

989 headers=headers, 

990 json=json_data, 

991 ) 

992 if r.status_code == 429: 

993 retry = float(r.headers.get('Retry-After', '1')) 

994 time.sleep(retry) 

995 continue 

996 r.raise_for_status() 

997 data = r.json() 

998 if 'data' in data: 

999 return [elem['embedding'] for elem in data['data']] 

1000 else: 

1001 raise ValueError("Unexpected Azure OpenAI embeddings response: missing 'data' key") 

1002 raise Exception('Azure OpenAI embedding request failed: max retries (429) exceeded') 

1003 

1004 

1005async def agenerate_azure_openai_batch_embeddings( 

1006 model: str, 

1007 texts: list[str], 

1008 url: str, 

1009 key: str = '', 

1010 version: str = '', 

1011 prefix: str = None, 

1012 user: UserModel = None, 

1013) -> list[list[float]]: 

1014 log.debug('agenerate_azure_openai_batch_embeddings:deployment %s batch size: %s', model, len(texts)) 

1015 form_data = {'input': texts} 

1016 if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): 

1017 form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix 

1018 

1019 full_url = f'{url}/openai/deployments/{model}/embeddings?api-version={version}' 

1020 

1021 headers = { 

1022 'Content-Type': 'application/json', 

1023 'api-key': key, 

1024 } 

1025 if ENABLE_FORWARD_USER_INFO_HEADERS and user: 

1026 headers = include_user_info_headers(headers, user) 

1027 

1028 async with aiohttp.ClientSession( 

1029 trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) 

1030 ) as session: 

1031 async with session.post( 

1032 full_url, 

1033 headers=headers, 

1034 json=form_data, 

1035 ssl=AIOHTTP_CLIENT_SESSION_SSL, 

1036 ) as r: 

1037 r.raise_for_status() 

1038 data = await r.json() 

1039 if 'data' in data: 

1040 return [item['embedding'] for item in data['data']] 

1041 else: 

1042 raise ValueError("Unexpected Azure OpenAI embeddings response: missing 'data' key") 

1043 

1044 

1045def generate_ollama_batch_embeddings( 

1046 model: str, 

1047 texts: list[str], 

1048 url: str, 

1049 key: str = '', 

1050 prefix: str = None, 

1051 user: UserModel = None, 

1052) -> list[list[float]]: 

1053 log.debug('generate_ollama_batch_embeddings:model %s batch size: %s', model, len(texts)) 

1054 json_data = {'input': texts, 'model': model, 'truncate': True} 

1055 if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): 

1056 json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix 

1057 

1058 headers = get_json_bearer_headers(key) 

1059 if ENABLE_FORWARD_USER_INFO_HEADERS and user: 

1060 headers = include_user_info_headers(headers, user) 

1061 

1062 r = requests.post( 

1063 f'{url}/api/embed', 

1064 headers=headers, 

1065 json=json_data, 

1066 ) 

1067 if r.status_code != 200: 

1068 error_detail = r.json().get('error', r.text) 

1069 raise Exception(f'Ollama embed error ({r.status_code}): {error_detail}') 

1070 data = r.json() 

1071 

1072 if 'embeddings' in data: 

1073 return data['embeddings'] 

1074 else: 

1075 raise ValueError("Unexpected Ollama embeddings response: missing 'embeddings' key") 

1076 

1077 

1078async def agenerate_ollama_batch_embeddings( 

1079 model: str, 

1080 texts: list[str], 

1081 url: str, 

1082 key: str = '', 

1083 prefix: str = None, 

1084 user: UserModel = None, 

1085) -> list[list[float]]: 

1086 log.debug('agenerate_ollama_batch_embeddings:model %s batch size: %s', model, len(texts)) 

1087 form_data = {'input': texts, 'model': model, 'truncate': True} 

1088 if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str): 

1089 form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix 

1090 

1091 headers = get_json_bearer_headers(key) 

1092 if ENABLE_FORWARD_USER_INFO_HEADERS and user: 

1093 headers = include_user_info_headers(headers, user) 

1094 

1095 async with aiohttp.ClientSession( 

1096 trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) 

1097 ) as session: 

1098 async with session.post( 

1099 f'{url}/api/embed', 

1100 headers=headers, 

1101 json=form_data, 

1102 ssl=AIOHTTP_CLIENT_SESSION_SSL, 

1103 ) as r: 

1104 if r.status != 200: 

1105 error_data = await r.json() 

1106 error_detail = error_data.get('error', str(error_data)) 

1107 raise Exception(f'Ollama embed error ({r.status}): {error_detail}') 

1108 data = await r.json() 

1109 if 'embeddings' in data: 

1110 return data['embeddings'] 

1111 else: 

1112 raise ValueError("Unexpected Ollama embeddings response: missing 'embeddings' key") 

1113 

1114 

1115def get_embedding_function( 

1116 embedding_engine, 

1117 embedding_model, 

1118 embedding_function, 

1119 url, 

1120 key, 

1121 embedding_batch_size, 

1122 azure_api_version=None, 

1123 enable_async=True, 

1124 concurrent_requests=0, 

1125) -> Awaitable: 

1126 if embedding_engine == '': 

1127 # Sentence transformers: CPU-bound sync operation 

1128 async def async_embedding_function(query, prefix=None, user=None): 

1129 if USE_SLIM: 1129 ↛ 1130line 1129 didn't jump to line 1130 because the condition on line 1129 was never true

1130 raise HTTPException(503, 'Configure an external embedding engine (openai, ollama, azure_openai).') 

1131 # Deferred so a missing local model degrades RAG instead of crashing boot. 

1132 if embedding_function is None: 

1133 raise ValueError( 

1134 'No embedding model is loaded. Set RAG_EMBEDDING_MODEL to a valid ' 

1135 'SentenceTransformer model name, or configure an external ' 

1136 'RAG_EMBEDDING_ENGINE (ollama, openai, azure_openai).' 

1137 ) 

1138 

1139 def encode(): 

1140 with MPS_INFERENCE_LOCK: 

1141 return embedding_function.encode( 

1142 query, 

1143 batch_size=int(embedding_batch_size), 

1144 **({'prompt': prefix} if prefix else {}), 

1145 ).tolist() 

1146 

1147 return await asyncio.to_thread(encode) 

1148 

1149 return async_embedding_function 

1150 elif embedding_engine in ['ollama', 'openai', 'azure_openai']: 1150 ↛ 1151line 1150 didn't jump to line 1151 because the condition on line 1150 was never true

1151 embedding_function = lambda query, prefix=None, user=None: generate_embeddings( 

1152 engine=embedding_engine, 

1153 model=embedding_model, 

1154 text=query, 

1155 prefix=prefix, 

1156 url=url, 

1157 key=key, 

1158 user=user, 

1159 azure_api_version=azure_api_version, 

1160 ) 

1161 

1162 async def async_embedding_function(query, prefix=None, user=None): 

1163 if isinstance(query, list): 

1164 # Create batches 

1165 batches = [query[i : i + embedding_batch_size] for i in range(0, len(query), embedding_batch_size)] 

1166 

1167 if enable_async: 

1168 log.debug('generate_multiple_async: Processing %s batches in parallel', len(batches)) 

1169 # Use semaphore to limit concurrent embedding API requests 

1170 # 0 = unlimited (no semaphore) 

1171 if concurrent_requests: 

1172 semaphore = asyncio.Semaphore(concurrent_requests) 

1173 

1174 async def generate_batch_with_semaphore(batch): 

1175 async with semaphore: 

1176 return await embedding_function(batch, prefix=prefix, user=user) 

1177 

1178 tasks = [generate_batch_with_semaphore(batch) for batch in batches] 

1179 else: 

1180 tasks = [embedding_function(batch, prefix=prefix, user=user) for batch in batches] 

1181 batch_results = await asyncio.gather(*tasks) 

1182 else: 

1183 log.debug('generate_multiple_async: Processing %s batches sequentially', len(batches)) 

1184 batch_results = [] 

1185 for batch in batches: 

1186 batch_results.append(await embedding_function(batch, prefix=prefix, user=user)) 

1187 

1188 # Flatten results — raise if any batch failed 

1189 embeddings = [] 

1190 for i, batch_embeddings in enumerate(batch_results): 

1191 if batch_embeddings is None: 

1192 raise Exception(f'Embedding generation failed for batch {i + 1}/{len(batches)}') 

1193 embeddings.extend(batch_embeddings) 

1194 

1195 log.debug( 

1196 'generate_multiple_async: Generated %s embeddings from %s parallel batches', 

1197 len(embeddings), 

1198 len(batches), 

1199 ) 

1200 return embeddings 

1201 else: 

1202 return await embedding_function(query, prefix, user) 

1203 

1204 return async_embedding_function 

1205 else: 

1206 raise ValueError(f'Unknown embedding engine: {embedding_engine}') 

1207 

1208 

1209async def generate_embeddings( 

1210 engine: str, 

1211 model: str, 

1212 text: Union[str, list[str]], 

1213 prefix: Union[str, None] = None, 

1214 **kwargs, 

1215): 

1216 url = kwargs.get('url', '') 

1217 key = kwargs.get('key', '') 

1218 user = kwargs.get('user') 

1219 

1220 if prefix is not None and RAG_EMBEDDING_PREFIX_FIELD_NAME is None: 

1221 if isinstance(text, list): 

1222 text = [f'{prefix}{text_element}' for text_element in text] 

1223 else: 

1224 text = f'{prefix}{text}' 

1225 

1226 if engine == 'ollama': 

1227 embeddings = await agenerate_ollama_batch_embeddings( 

1228 **{ 

1229 'model': model, 

1230 'texts': text if isinstance(text, list) else [text], 

1231 'url': url, 

1232 'key': key, 

1233 'prefix': prefix, 

1234 'user': user, 

1235 } 

1236 ) 

1237 if embeddings is None: 

1238 return None 

1239 return embeddings[0] if isinstance(text, str) else embeddings 

1240 elif engine == 'openai': 

1241 embeddings = await agenerate_openai_batch_embeddings( 

1242 model, text if isinstance(text, list) else [text], url, key, prefix, user 

1243 ) 

1244 if embeddings is None: 

1245 return None 

1246 return embeddings[0] if isinstance(text, str) else embeddings 

1247 elif engine == 'azure_openai': 

1248 azure_api_version = kwargs.get('azure_api_version', '') 

1249 embeddings = await agenerate_azure_openai_batch_embeddings( 

1250 model, 

1251 text if isinstance(text, list) else [text], 

1252 url, 

1253 key, 

1254 azure_api_version, 

1255 prefix, 

1256 user, 

1257 ) 

1258 if embeddings is None: 

1259 return None 

1260 return embeddings[0] if isinstance(text, str) else embeddings 

1261 

1262 

1263def get_reranking_function(reranking_engine, reranking_model, reranking_function, reranking_batch_size=32): 

1264 if USE_SLIM and reranking_model and reranking_engine != 'external': 1264 ↛ 1266line 1264 didn't jump to line 1266 because the condition on line 1264 was never true

1265 

1266 def unavailable(query, documents, user=None): 

1267 raise HTTPException( 

1268 503, 'Configure an external reranker, or clear the reranking model to use cosine scoring.' 

1269 ) 

1270 

1271 return unavailable 

1272 if reranking_function is None: 1272 ↛ 1274line 1272 didn't jump to line 1274 because the condition on line 1272 was always true

1273 return None 

1274 if reranking_engine == 'external': 

1275 return lambda query, documents, user=None: reranking_function.predict( 

1276 [(query, doc.page_content) for doc in documents], user=user 

1277 ) 

1278 else: 

1279 

1280 def predict(query, documents, user=None): 

1281 with MPS_INFERENCE_LOCK: 

1282 return reranking_function.predict( 

1283 [(query, doc.page_content) for doc in documents], batch_size=int(reranking_batch_size) 

1284 ) 

1285 

1286 return predict 

1287 

1288 

1289# UUIDs, SHA-256 digests, and prefixed variants thereof all fit [A-Za-z0-9_-]. 

1290# Anything else cannot be a real Open WebUI collection and could break out of 

1291# a Milvus expression literal. 

1292_SAFE_COLLECTION_NAME_RE = re.compile(r'^[A-Za-z0-9_-]{1,255}$') 

1293 

1294 

1295def _is_safe_collection_name(name: str) -> bool: 

1296 return isinstance(name, str) and bool(_SAFE_COLLECTION_NAME_RE.match(name)) 

1297 

1298 

1299async def filter_accessible_collections( 

1300 collection_names: set[str], 

1301 user: UserModel, 

1302 access_type: str = 'read', 

1303) -> set[str]: 

1304 """ 

1305 Return only the collection names the user is allowed to access. 

1306 Admins bypass all checks. For non-admins the policy is: 

1307 

1308 - any name with characters outside [A-Za-z0-9_-] → rejected 

1309 - file-* → validated via has_access_to_file 

1310 - user-memory-* → must match user's own memory collection 

1311 - web-search-* → ephemeral per-query collections, owner-bound to web-search-{user.id}-* 

1312 - knowledge-bases → always denied (system meta-collection) 

1313 - everything else → if the name matches a knowledge base, validated 

1314 via Knowledges.check_access_by_user_id; if no 

1315 such KB exists, denied by default. When 

1316 ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS is True, 

1317 the name is treated as a legacy/ephemeral 

1318 collection and allowed. 

1319 """ 

1320 # Applied before the admin bypass — malformed names should never reach the vector store. 

1321 safe_names = {n for n in collection_names if _is_safe_collection_name(n)} 

1322 rejected = collection_names - safe_names 

1323 if rejected: 

1324 log.warning( 

1325 'filter_accessible_collections: rejected %d collection name(s) with unsafe characters (user_id=%s)', 

1326 len(rejected), 

1327 getattr(user, 'id', '<unknown>'), 

1328 ) 

1329 

1330 if user.role == 'admin': 1330 ↛ 1333line 1330 didn't jump to line 1333 because the condition on line 1330 was always true

1331 return safe_names 

1332 

1333 validated = set() 

1334 for name in safe_names: 

1335 if name == 'knowledge-bases': 

1336 # System meta-collection — never exposed to non-admins. 

1337 continue 

1338 elif name.startswith('file-'): 

1339 file_id = name[len('file-') :] 

1340 if await has_access_to_file(file_id=file_id, access_type=access_type, user=user): 

1341 validated.add(name) 

1342 elif name.startswith('user-memory-'): 

1343 if name == f'user-memory-{user.id}': 

1344 validated.add(name) 

1345 elif name.startswith('web-search-'): 

1346 # Ephemeral per-query collections, owner-bound: process_web_search mints 

1347 # them as web-search-{user.id}-<hash>, so only the creator may read/write. 

1348 if name.startswith(f'web-search-{user.id}-'): 

1349 validated.add(name) 

1350 else: 

1351 # May be a knowledge-base ID or a legacy/ephemeral collection. 

1352 # If it IS a KB, enforce access control. If no such KB 

1353 # exists, the behaviour depends on 

1354 # ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS: 

1355 # False (default) — deny (closes the unscoped namespace) 

1356 # True — allow (preserves legacy behaviour) 

1357 if await Knowledges.check_access_by_user_id(name, user.id, permission=access_type): 

1358 validated.add(name) 

1359 elif ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS and not await Knowledges.get_knowledge_by_id(name): 

1360 # Not a KB at all — legacy/ephemeral collection, allow 

1361 validated.add(name) 

1362 return validated 

1363 

1364 

1365def filter_source_metadata(metadata: dict) -> dict: 

1366 """Keep only the chunk metadata keys the operator allowed the model to see.""" 

1367 return {key: metadata[key] for key in RAG_SOURCE_METADATA_KEYS if metadata.get(key) is not None} 

1368 

1369 

1370async def get_sources_from_items( 

1371 request, 

1372 items, 

1373 queries, 

1374 embedding_function, 

1375 k, 

1376 reranking_function, 

1377 k_reranker, 

1378 r, 

1379 hybrid_bm25_weight, 

1380 hybrid_search, 

1381 full_context=False, 

1382 user: UserModel | None = None, 

1383): 

1384 log.debug('items: %s %s %s %s %s', items, queries, embedding_function, reranking_function, full_context) 

1385 

1386 bypass_embedding_and_retrieval = await Config.get('rag.bypass_embedding_and_retrieval') 

1387 extracted_collections = [] 

1388 query_results = [] 

1389 folder_items = set() 

1390 expanded_folders = set() 

1391 

1392 items = list(items) 

1393 for item in items: 

1394 if item.get('type') != 'folder' or not user: 

1395 continue 

1396 folder_id = item.get('id') 

1397 if not folder_id or folder_id in expanded_folders: 

1398 continue 

1399 expanded_folders.add(folder_id) 

1400 

1401 folder = await Folders.get_folder_by_id(folder_id) 

1402 if folder and (user.role == 'admin' or await has_folder_access(user.id, folder, 'read', db=None)): 

1403 files = await get_owner_accessible_folder_files(folder) 

1404 folder_items.update((entry.get('type'), entry.get('id')) for entry in files if isinstance(entry, dict)) 

1405 items.extend(files) 

1406 

1407 for item in items: 

1408 query_result = None 

1409 collection_names = [] 

1410 

1411 if item.get('type') == 'text': 

1412 # Raw Text 

1413 # Used during temporary chat file uploads or web page & youtube attachements 

1414 

1415 if item.get('context') == 'full': 

1416 if item.get('file'): 

1417 # if item has file data, use it 

1418 query_result = { 

1419 'documents': [[item.get('file', {}).get('data', {}).get('content')]], 

1420 'metadatas': [[item.get('file', {}).get('meta', {})]], 

1421 } 

1422 

1423 if query_result is None: 

1424 # Fallback 

1425 if item.get('collection_name'): 

1426 # If item has a collection name, use it 

1427 collection_names.append(item.get('collection_name')) 

1428 elif item.get('file'): 

1429 # If item has file data, use it 

1430 query_result = { 

1431 'documents': [[item.get('file', {}).get('data', {}).get('content')]], 

1432 'metadatas': [[item.get('file', {}).get('meta', {})]], 

1433 } 

1434 else: 

1435 # Fallback to item content 

1436 query_result = { 

1437 'documents': [[item.get('content')]], 

1438 'metadatas': [[{'file_id': item.get('id'), 'name': item.get('name')}]], 

1439 } 

1440 

1441 elif item.get('type') == 'note': 

1442 # Note Attached 

1443 note = await Notes.get_note_by_id(item.get('id')) 

1444 

1445 if note and ( 

1446 user.role == 'admin' 

1447 or note.user_id == user.id 

1448 or await AccessGrants.has_access( 

1449 user_id=user.id, 

1450 resource_type='note', 

1451 resource_id=note.id, 

1452 permission='read', 

1453 ) 

1454 ): 

1455 # User has access to the note 

1456 query_result = { 

1457 'documents': [[note.data.get('content', {}).get('md', '')]], 

1458 'metadatas': [[{'file_id': note.id, 'name': note.title}]], 

1459 } 

1460 

1461 elif item.get('type') == 'chat': 

1462 # Chat Attached 

1463 chat = await Chats.get_chat_by_id(item.get('id')) 

1464 has_read_access = bool(chat and (user.role == 'admin' or chat.user_id == user.id)) 

1465 

1466 if chat and not has_read_access: 

1467 has_read_access = await AccessGrants.has_access( 

1468 user_id=user.id, 

1469 resource_type='shared_chat', 

1470 resource_id=chat.id, 

1471 permission='read', 

1472 ) 

1473 

1474 if chat and not has_read_access and chat.folder_id: 

1475 folder = await Folders.get_folder_by_id(chat.folder_id) 

1476 has_read_access = folder and await has_folder_access(user.id, folder, 'read', db=None) 

1477 

1478 if has_read_access: 

1479 messages_map = chat.chat.get('history', {}).get('messages', {}) 

1480 message_id = chat.chat.get('history', {}).get('currentId') 

1481 

1482 if messages_map and message_id: 

1483 # Reconstruct the message list in order 

1484 message_list = get_message_list(messages_map, message_id) 

1485 message_history = '\n'.join( 

1486 [ 

1487 f'#### {m.get("role", "user").capitalize()}\n{get_content_from_message(m) or ""}\n' 

1488 for m in message_list 

1489 ] 

1490 ) 

1491 

1492 # User has access to the chat 

1493 query_result = { 

1494 'documents': [[message_history]], 

1495 'metadatas': [[{'file_id': chat.id, 'name': chat.title}]], 

1496 } 

1497 

1498 elif item.get('type') == 'url': 

1499 content, docs = await get_content_from_url(request, item.get('url')) 

1500 if docs: 

1501 query_result = { 

1502 'documents': [[content]], 

1503 'metadatas': [[{'url': item.get('url'), 'name': item.get('url')}]], 

1504 } 

1505 elif item.get('type') == 'file': 

1506 if item.get('context') == 'full' or bypass_embedding_and_retrieval: 

1507 if item.get('file', {}).get('data', {}).get('content', ''): 

1508 # Manual Full Mode Toggle 

1509 # Used from chat file modal, we can assume that the file content will be available from item.get("file").get("data", {}).get("content") 

1510 query_result = { 

1511 'documents': [[item.get('file', {}).get('data', {}).get('content', '')]], 

1512 'metadatas': [ 

1513 [ 

1514 { 

1515 'file_id': item.get('id'), 

1516 'name': item.get('name'), 

1517 **item.get('file').get('data', {}).get('metadata', {}), 

1518 } 

1519 ] 

1520 ], 

1521 } 

1522 elif item.get('id'): 

1523 file_object = await Files.get_file_by_id(item.get('id')) 

1524 if file_object and ( 

1525 user.role == 'admin' 

1526 or file_object.user_id == user.id 

1527 or await has_access_to_file(item.get('id'), 'read', user) 

1528 or ('file', item.get('id')) in folder_items 

1529 ): 

1530 query_result = { 

1531 'documents': [[file_object.data.get('content', '')]], 

1532 'metadatas': [ 

1533 [ 

1534 { 

1535 'file_id': item.get('id'), 

1536 'name': file_object.filename, 

1537 'source': file_object.filename, 

1538 } 

1539 ] 

1540 ], 

1541 } 

1542 else: 

1543 # Chunked-retrieval fallback — verify read access before 

1544 # exposing the file's vector collection (same posture as the 

1545 # full-context branch above). 

1546 file_id = item.get('id') 

1547 if file_id: 

1548 if BYPASS_RETRIEVAL_ACCESS_CONTROL: 

1549 if item.get('legacy'): 

1550 collection_names.append(f'{file_id}') 

1551 else: 

1552 collection_names.append(f'file-{file_id}') 

1553 else: 

1554 file_object = await Files.get_file_by_id(file_id) 

1555 if file_object and ( 

1556 user.role == 'admin' 

1557 or file_object.user_id == user.id 

1558 or await has_access_to_file(file_id, 'read', user) 

1559 or ('file', file_id) in folder_items 

1560 ): 

1561 if item.get('legacy'): 

1562 collection_names.append(f'{file_id}') 

1563 else: 

1564 collection_names.append(f'file-{file_id}') 

1565 

1566 elif item.get('type') == 'collection': 

1567 # Manual Full Mode Toggle for Collection 

1568 knowledge_base = await Knowledges.get_knowledge_by_id(item.get('id')) 

1569 

1570 if knowledge_base and ( 

1571 user.role == 'admin' 

1572 or knowledge_base.user_id == user.id 

1573 or await AccessGrants.has_access( 

1574 user_id=user.id, 

1575 resource_type='knowledge', 

1576 resource_id=knowledge_base.id, 

1577 permission='read', 

1578 ) 

1579 or ('collection', item.get('id')) in folder_items 

1580 ): 

1581 if (knowledge_base.meta or {}).get('source') == 'external': 

1582 query_result = await retrieve_external_knowledge( 

1583 request, 

1584 knowledge_base, 

1585 queries=queries, 

1586 count=k, 

1587 user=user, 

1588 ) 

1589 extracted_collections.append(knowledge_base.id) 

1590 

1591 else: 

1592 if item.get('context') == 'full' or bypass_embedding_and_retrieval: 

1593 if knowledge_base and ( 

1594 user.role == 'admin' 

1595 or knowledge_base.user_id == user.id 

1596 or await AccessGrants.has_access( 

1597 user_id=user.id, 

1598 resource_type='knowledge', 

1599 resource_id=knowledge_base.id, 

1600 permission='read', 

1601 ) 

1602 or ('collection', item.get('id')) in folder_items 

1603 ): 

1604 files = await Knowledges.get_files_by_id(knowledge_base.id) 

1605 

1606 documents = [] 

1607 metadatas = [] 

1608 for file in files: 

1609 documents.append(file.data.get('content', '')) 

1610 metadatas.append( 

1611 { 

1612 'file_id': file.id, 

1613 'name': file.filename, 

1614 'source': file.filename, 

1615 } 

1616 ) 

1617 

1618 query_result = { 

1619 'documents': [documents], 

1620 'metadatas': [metadatas], 

1621 } 

1622 else: 

1623 if item.get('legacy'): 

1624 if BYPASS_RETRIEVAL_ACCESS_CONTROL: 

1625 collection_names = item.get('collection_names', []) 

1626 else: 

1627 # Legacy KB: item.collection_names is client-supplied. 

1628 # Validate against the KB's actual files to prevent 

1629 # cross-tenant collection name substitution. 

1630 files = await Knowledges.get_files_by_id(knowledge_base.id) 

1631 owned_names = {f'file-{f.id}' for f in files} 

1632 owned_names.add(knowledge_base.id) 

1633 valid_names = [n for n in (item.get('collection_names') or []) if n in owned_names] 

1634 collection_names = valid_names if valid_names else [knowledge_base.id] 

1635 else: 

1636 collection_names.append(item['id']) 

1637 

1638 elif item.get('docs'): 

1639 # BYPASS_WEB_SEARCH_EMBEDDING_AND_RETRIEVAL 

1640 query_result = { 

1641 'documents': [[doc.get('content') for doc in item.get('docs')]], 

1642 'metadatas': [[doc.get('metadata') for doc in item.get('docs')]], 

1643 } 

1644 elif item.get('type') == 'web_search' and item.get('collection_name'): 

1645 # Trusted server-generated collection; authorized by 

1646 # filter_accessible_collections below (allowlists web-search-*). 

1647 collection_names.append(item['collection_name']) 

1648 elif item.get('collection_name'): 

1649 if BYPASS_RETRIEVAL_ACCESS_CONTROL: 

1650 collection_names.append(item['collection_name']) 

1651 else: 

1652 log.debug( 

1653 "get_sources_from_items: ignoring untrusted direct collection_name '%s' on item without type", 

1654 item.get('collection_name'), 

1655 ) 

1656 elif item.get('collection_names'): 

1657 if BYPASS_RETRIEVAL_ACCESS_CONTROL: 

1658 collection_names.extend(item['collection_names']) 

1659 else: 

1660 log.debug( 

1661 'get_sources_from_items: ignoring untrusted direct collection_names on item without type', 

1662 ) 

1663 

1664 # If query_result is None 

1665 # Fallback to collection names and vector search the collections 

1666 if query_result is None and collection_names: 

1667 collection_names = set(collection_names).difference(extracted_collections) 

1668 if not collection_names: 

1669 log.debug('skipping %s as it has already been extracted', item) 

1670 continue 

1671 

1672 # Filter out collections the user cannot read 

1673 if user and (item.get('type'), item.get('id')) not in folder_items: 

1674 collection_names = await filter_accessible_collections(collection_names, user) 

1675 if not collection_names: 

1676 log.debug('access denied for all collections in item %s', item) 

1677 continue 

1678 

1679 try: 

1680 if full_context: 

1681 # Sync helper makes blocking VECTOR_DB_CLIENT calls; 

1682 # offload so the async caller's event loop stays free. 

1683 query_result = await asyncio.to_thread(get_all_items_from_collections, collection_names) 

1684 else: 

1685 query_result = await query_collection( 

1686 request, 

1687 collection_names=collection_names, 

1688 queries=queries, 

1689 embedding_function=embedding_function, 

1690 k=k, 

1691 ) 

1692 except Exception as e: 

1693 log.exception(e) 

1694 

1695 extracted_collections.extend(collection_names) 

1696 

1697 if query_result: 

1698 if 'data' in item: 

1699 del item['data'] 

1700 query_results.append({**query_result, 'file': item}) 

1701 

1702 sources = [] 

1703 for query_result in query_results: 

1704 try: 

1705 if 'documents' in query_result: 

1706 if 'metadatas' in query_result: 

1707 source = { 

1708 'source': query_result['file'], 

1709 'document': query_result['documents'][0], 

1710 'metadata': query_result['metadatas'][0], 

1711 } 

1712 if 'distances' in query_result and query_result['distances']: 

1713 source['distances'] = query_result['distances'][0] 

1714 

1715 sources.append(source) 

1716 except Exception as e: 

1717 log.exception(e) 

1718 return sources 

1719 

1720 

1721def get_model_path(model: str, update_model: bool = False): 

1722 from huggingface_hub import snapshot_download 

1723 

1724 # Construct huggingface_hub kwargs with local_files_only to return the snapshot path 

1725 cache_dir = os.getenv('SENTENCE_TRANSFORMERS_HOME') 

1726 

1727 local_files_only = not update_model 

1728 

1729 if OFFLINE_MODE: 1729 ↛ 1732line 1729 didn't jump to line 1732 because the condition on line 1729 was always true

1730 local_files_only = True 

1731 

1732 snapshot_kwargs = { 

1733 'cache_dir': cache_dir, 

1734 'local_files_only': local_files_only, 

1735 } 

1736 

1737 log.debug('model: %s', model) 

1738 log.debug('snapshot_kwargs: %s', snapshot_kwargs) 

1739 

1740 # Inspiration from upstream sentence_transformers 

1741 if os.path.exists(model) or ('\\' in model or model.count('/') > 1) and local_files_only: 1741 ↛ 1743line 1741 didn't jump to line 1743 because the condition on line 1741 was never true

1742 # If fully qualified path exists, return input, else set repo_id 

1743 return model 

1744 elif '/' not in model: 

1745 # Set valid repo_id for model short-name 

1746 model = 'sentence-transformers' + '/' + model 

1747 

1748 snapshot_kwargs['repo_id'] = model 

1749 

1750 # Attempt to query the huggingface_hub library to determine the local path and/or to update 

1751 try: 

1752 model_repo_path = snapshot_download(**snapshot_kwargs) 

1753 log.debug('model_repo_path: %s', model_repo_path) 

1754 return model_repo_path 

1755 except Exception as e: 

1756 log.exception(f'Cannot determine model snapshot path: {e}') 

1757 if OFFLINE_MODE: 1757 ↛ 1759line 1757 didn't jump to line 1759 because the condition on line 1757 was always true

1758 raise 

1759 return model 

1760 

1761 

1762import operator 

1763from typing import Optional, Sequence 

1764 

1765from langchain_core.callbacks import Callbacks 

1766from langchain_core.documents import BaseDocumentCompressor, Document 

1767 

1768 

1769def cosine_similarity(query, documents) -> np.ndarray: 

1770 """Score one query against documents without loading a model runtime.""" 

1771 if len(documents) == 0: 

1772 return np.array([], dtype=float) 

1773 query = np.asarray(query, dtype=float).reshape(-1) 

1774 documents = np.asarray(documents, dtype=float) 

1775 query = query / max(np.linalg.norm(query), 1e-12) 

1776 documents = documents / np.maximum(np.linalg.norm(documents, axis=1, keepdims=True), 1e-12) 

1777 return documents @ query 

1778 

1779 

1780class RerankCompressor(BaseDocumentCompressor): 

1781 embedding_function: Any 

1782 top_n: int 

1783 reranking_function: Any 

1784 r_score: float 

1785 

1786 class Config: 

1787 extra = 'forbid' 

1788 arbitrary_types_allowed = True 

1789 

1790 def compress_documents( 

1791 self, 

1792 documents: Sequence[Document], 

1793 query: str, 

1794 callbacks: Callbacks | None = None, 

1795 ) -> Sequence[Document]: 

1796 """Compress retrieved documents given the query context. 

1797 

1798 Args: 

1799 documents: The retrieved documents. 

1800 query: The query context. 

1801 callbacks: Optional callbacks to run during compression. 

1802 

1803 Returns: 

1804 The compressed documents. 

1805 

1806 """ 

1807 return [] 

1808 

1809 async def acompress_documents( 

1810 self, 

1811 documents: Sequence[Document], 

1812 query: str, 

1813 callbacks: Callbacks | None = None, 

1814 ) -> Sequence[Document]: 

1815 if not documents: 

1816 return [] 

1817 reranking = self.reranking_function is not None 

1818 

1819 scores = None 

1820 if reranking: 

1821 scores = await asyncio.to_thread(self.reranking_function, query, documents) 

1822 else: 

1823 query_embedding = await self.embedding_function(query, RAG_EMBEDDING_QUERY_PREFIX) 

1824 doc_texts = [doc.page_content for doc in documents] 

1825 document_embedding = await self.embedding_function(doc_texts, RAG_EMBEDDING_CONTENT_PREFIX) 

1826 scores = cosine_similarity(query_embedding, document_embedding) 

1827 

1828 if scores is not None: 

1829 docs_with_scores = list( 

1830 zip( 

1831 documents, 

1832 scores.tolist() if not isinstance(scores, list) else scores, 

1833 ) 

1834 ) 

1835 if self.r_score: 

1836 docs_with_scores = [(d, s) for d, s in docs_with_scores if s >= self.r_score] 

1837 

1838 result = sorted(docs_with_scores, key=operator.itemgetter(1), reverse=True) 

1839 final_results = [] 

1840 for doc, doc_score in result[: self.top_n]: 

1841 metadata = doc.metadata 

1842 metadata['score'] = doc_score 

1843 doc = Document( 

1844 page_content=doc.page_content, 

1845 metadata=metadata, 

1846 ) 

1847 final_results.append(doc) 

1848 return final_results 

1849 else: 

1850 log.warning('No valid scores found, check your reranking function. Returning original documents.') 

1851 return documents