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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1from __future__ import annotations
3import asyncio
4import hashlib
5import logging
6import os
7import re
8import time
9from typing import Awaitable, Optional, Union
10from urllib.parse import quote
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
59log = logging.getLogger(__name__)
62from typing import Any
64from langchain_core.callbacks import CallbackManagerForRetrieverRun
65from langchain_core.retrievers import BaseRetriever
68class BM25Retriever(BaseRetriever):
69 docs: list[Document]
70 vectorizer: Any
71 k: int
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)
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
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}
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()}
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 )
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
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 )
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
183 content_type = response.headers.get('Content-Type', '').split(';')[0].strip()
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 ''
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('"\'')
195 if not filename or '.' not in filename:
196 ext = mimetypes.guess_extension(content_type) or ''
197 filename = f'download{ext}'
199 suffix = '.' + filename.split('.')[-1].lower() if '.' in filename else ''
201 max_size = loader_config.get('file_max_size')
202 max_bytes = int(max_size) * 1024 * 1024 if max_size else 0
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)
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)
225TEXT_APPLICATION_CONTENT_TYPES = {
226 'application/javascript',
227 'application/json',
228 'application/xml',
229 'application/x-javascript',
230}
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'))
245async def get_content_from_url(request, url: str) -> str:
246 loader_config = await get_loader_config()
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)
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
257 # Validate URL before making any request (blocks private IPs, non-HTTP, filter list)
258 validate_url(url)
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
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
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
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()
303CHUNK_HASH_KEY = '_chunk_hash'
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()
311class VectorSearchRetriever(BaseRetriever):
312 collection_name: Any
313 embedding_function: Any
314 top_k: int
316 def _get_relevant_documents(self, query: str, *, run_manager: CallbackManagerForRetrieverRun) -> list[Document]:
317 """Get documents relevant to a query.
319 Args:
320 query: String to find relevant documents for.
321 run_manager: The callback handler to use.
323 Returns:
324 List of relevant documents.
325 """
326 return []
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 )
341 return _search_result_to_documents(result)
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 )
352 if result:
353 log.info('query_doc:result %s %s', result.ids, result.metadatas)
355 return result
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)
363 if result:
364 log.info('query_doc:result %s %s', result.ids, result.metadatas)
366 return result
367 except Exception as e:
368 log.exception(f'Error getting doc {collection_name}: {e}')
369 raise e
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]
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}')
384 # Add title if available
385 if metadata.get('title'):
386 metadata_parts.append(f'Title: {metadata["title"]}')
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}')
393 # Add source URL/path if available
394 if metadata.get('source'):
395 metadata_parts.append(f'Source: {metadata["source"]}')
397 # Add snippet for web search results
398 if metadata.get('snippet'):
399 metadata_parts.append(f'Snippet: {metadata["snippet"]}')
401 enriched_texts.append(' '.join(metadata_parts))
403 return enriched_texts
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 []
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
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))
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
444 query_vectors = []
445 if hybrid_bm25_weight < 1:
446 query_vectors = [await embedding_function(query, RAG_EMBEDDING_QUERY_PREFIX)]
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
458 documents = _search_result_to_documents(result)
459 if not documents:
460 return {'distances': [[]], 'documents': [[]], 'metadatas': [[]]}
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)
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]
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]
478 if sorted_items:
479 distances, documents, metadatas = map(list, zip(*sorted_items))
480 else:
481 distances, documents, metadatas = [], [], []
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
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
520 if collection_result is None:
521 collection_result = await ASYNC_VECTOR_DB_CLIENT.get(collection_name=collection_name)
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': []}
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': []}
537 log.debug('query_doc_with_hybrid_search:doc %s', collection_name)
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 ]
545 bm25_texts = get_enriched_texts(collection_result) if enable_enriched_texts else original_texts
547 from rank_bm25 import BM25Okapi
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 )
555 vector_search_retriever = VectorSearchRetriever(
556 collection_name=collection_name,
557 embedding_function=embedding_function,
558 top_k=k,
559 )
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 )
581 compressor = RerankCompressor(
582 embedding_function=embedding_function,
583 top_n=k_reranker,
584 reranking_function=reranking_function,
585 r_score=r,
586 )
588 compression_retriever = ContextualCompressionRetriever(
589 base_compressor=compressor, base_retriever=ensemble_retriever
590 )
592 result = await compression_retriever.ainvoke(query)
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]
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]
603 if sorted_items:
604 distances, documents, metadatas = map(list, zip(*sorted_items))
605 else:
606 distances, documents, metadatas = [], [], []
608 result = {
609 'distances': [distances],
610 'documents': [documents],
611 'metadatas': [metadatas],
612 }
614 log.info('query_doc_with_hybrid_search:result %s %s', result['metadatas'], result['distances'])
615 return result
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 = []
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])
629 # Create the output dictionary
630 result = {
631 'documents': [combined_documents],
632 'metadatas': [combined_metadatas],
633 'ids': [combined_ids],
634 }
636 return result
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
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
651 distances = data['distances'][0]
652 documents = data['documents'][0]
653 metadatas = data['metadatas'][0]
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)
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
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)
667 combined = list(combined.values())
668 # Sort the list based on distances
669 combined.sort(key=lambda x: x[0], reverse=True)
671 # Slice to keep only the top k elements
672 sorted_distances, sorted_documents, sorted_metadatas = zip(*combined[:k]) if combined else ([], [], [])
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 }
682def get_all_items_from_collections(collection_names: list[str]) -> dict:
683 results = []
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
696 return merge_get_results(results)
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)
735 results = []
736 last_error = None
737 failed_collection_names = set()
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
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': [[]]}
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))
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 )
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)
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 )
787 return merge_and_sort_query_results(results, k=k)
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()
805 if not enable_enriched_texts:
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
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)
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 )
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
843 collection_results = dict(await asyncio.gather(*(_fetch_collection(name) for name in collection_names)))
845 log.info('Starting hybrid search for %s queries in %s collections...', len(queries), len(collection_names))
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
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 ]
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])
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)
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 )
893 if failed_collection_names and not results:
894 raise Exception('Hybrid search failed for all collections. Using Non-hybrid search as fallback.')
896 return merge_and_sort_query_results(results, k=k)
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
912 headers = get_json_bearer_headers(key)
913 if ENABLE_FORWARD_USER_INFO_HEADERS and user:
914 headers = include_user_info_headers(headers, user)
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")
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
942 headers = get_json_bearer_headers(key)
943 if ENABLE_FORWARD_USER_INFO_HEADERS and user:
944 headers = include_user_info_headers(headers, user)
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")
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
977 url = f'{url}/openai/deployments/{model}/embeddings?api-version={version}'
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)
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')
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
1019 full_url = f'{url}/openai/deployments/{model}/embeddings?api-version={version}'
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)
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")
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
1058 headers = get_json_bearer_headers(key)
1059 if ENABLE_FORWARD_USER_INFO_HEADERS and user:
1060 headers = include_user_info_headers(headers, user)
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()
1072 if 'embeddings' in data:
1073 return data['embeddings']
1074 else:
1075 raise ValueError("Unexpected Ollama embeddings response: missing 'embeddings' key")
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
1091 headers = get_json_bearer_headers(key)
1092 if ENABLE_FORWARD_USER_INFO_HEADERS and user:
1093 headers = include_user_info_headers(headers, user)
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")
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 )
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()
1147 return await asyncio.to_thread(encode)
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 )
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)]
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)
1174 async def generate_batch_with_semaphore(batch):
1175 async with semaphore:
1176 return await embedding_function(batch, prefix=prefix, user=user)
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))
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)
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)
1204 return async_embedding_function
1205 else:
1206 raise ValueError(f'Unknown embedding engine: {embedding_engine}')
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')
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}'
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
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
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 )
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:
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 )
1286 return predict
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}$')
1295def _is_safe_collection_name(name: str) -> bool:
1296 return isinstance(name, str) and bool(_SAFE_COLLECTION_NAME_RE.match(name))
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:
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 )
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
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
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}
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)
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()
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)
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)
1407 for item in items:
1408 query_result = None
1409 collection_names = []
1411 if item.get('type') == 'text':
1412 # Raw Text
1413 # Used during temporary chat file uploads or web page & youtube attachements
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 }
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 }
1441 elif item.get('type') == 'note':
1442 # Note Attached
1443 note = await Notes.get_note_by_id(item.get('id'))
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 }
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))
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 )
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)
1478 if has_read_access:
1479 messages_map = chat.chat.get('history', {}).get('messages', {})
1480 message_id = chat.chat.get('history', {}).get('currentId')
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 )
1492 # User has access to the chat
1493 query_result = {
1494 'documents': [[message_history]],
1495 'metadatas': [[{'file_id': chat.id, 'name': chat.title}]],
1496 }
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}')
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'))
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)
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)
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 )
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'])
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 )
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
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
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)
1695 extracted_collections.extend(collection_names)
1697 if query_result:
1698 if 'data' in item:
1699 del item['data']
1700 query_results.append({**query_result, 'file': item})
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]
1715 sources.append(source)
1716 except Exception as e:
1717 log.exception(e)
1718 return sources
1721def get_model_path(model: str, update_model: bool = False):
1722 from huggingface_hub import snapshot_download
1724 # Construct huggingface_hub kwargs with local_files_only to return the snapshot path
1725 cache_dir = os.getenv('SENTENCE_TRANSFORMERS_HOME')
1727 local_files_only = not update_model
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
1732 snapshot_kwargs = {
1733 'cache_dir': cache_dir,
1734 'local_files_only': local_files_only,
1735 }
1737 log.debug('model: %s', model)
1738 log.debug('snapshot_kwargs: %s', snapshot_kwargs)
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
1748 snapshot_kwargs['repo_id'] = model
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
1762import operator
1763from typing import Optional, Sequence
1765from langchain_core.callbacks import Callbacks
1766from langchain_core.documents import BaseDocumentCompressor, Document
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
1780class RerankCompressor(BaseDocumentCompressor):
1781 embedding_function: Any
1782 top_n: int
1783 reranking_function: Any
1784 r_score: float
1786 class Config:
1787 extra = 'forbid'
1788 arbitrary_types_allowed = True
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.
1798 Args:
1799 documents: The retrieved documents.
1800 query: The query context.
1801 callbacks: Optional callbacks to run during compression.
1803 Returns:
1804 The compressed documents.
1806 """
1807 return []
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
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)
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]
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