From f7a6d517afc6691bcafcd80ff74569c92fb12e1f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=C3=89ole=20PASKALI?= Date: Sun, 14 Jun 2026 14:22:43 +0200 Subject: [PATCH] Add message-scoped file context option --- backend/open_webui/config.py | 6 + backend/open_webui/main.py | 2 + backend/open_webui/routers/retrieval.py | 8 + backend/open_webui/utils/middleware.py | 194 +++++++++++++++++- .../admin/Settings/Documents.svelte | 16 ++ 5 files changed, 218 insertions(+), 8 deletions(-) diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index 57f4712027..063a7096f2 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -1250,6 +1250,12 @@ RAG_FULL_CONTEXT = ConfigVar( os.getenv('RAG_FULL_CONTEXT', 'False').lower() == 'true', ) +RAG_MESSAGE_SCOPED_FILE_CONTEXT = ConfigVar( + 'RAG_MESSAGE_SCOPED_FILE_CONTEXT', + 'rag.message_scoped_file_context', + os.getenv('RAG_MESSAGE_SCOPED_FILE_CONTEXT', 'False').lower() == 'true', +) + RAG_FILE_MAX_COUNT = ConfigVar( 'RAG_FILE_MAX_COUNT', 'rag.file.max_count', diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 07d7583002..5de1754685 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -336,6 +336,7 @@ from open_webui.config import ( RAG_FILE_MAX_SIZE, RAG_FULL_CONTEXT, RAG_HYBRID_BM25_WEIGHT, + RAG_MESSAGE_SCOPED_FILE_CONTEXT, RAG_OLLAMA_API_KEY, RAG_OLLAMA_BASE_URL, RAG_OPENAI_API_BASE_URL, @@ -998,6 +999,7 @@ app.state.config.FILE_IMAGE_COMPRESSION_HEIGHT = FILE_IMAGE_COMPRESSION_HEIGHT app.state.config.RAG_FULL_CONTEXT = RAG_FULL_CONTEXT +app.state.config.RAG_MESSAGE_SCOPED_FILE_CONTEXT = RAG_MESSAGE_SCOPED_FILE_CONTEXT app.state.config.BYPASS_EMBEDDING_AND_RETRIEVAL = BYPASS_EMBEDDING_AND_RETRIEVAL app.state.config.ENABLE_RAG_HYBRID_SEARCH = ENABLE_RAG_HYBRID_SEARCH app.state.config.ENABLE_RAG_HYBRID_SEARCH_ENRICHED_TEXTS = ENABLE_RAG_HYBRID_SEARCH_ENRICHED_TEXTS diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index 70f6cf6309..d3126499d5 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -426,6 +426,7 @@ async def get_rag_config(request: Request, user=Depends(get_admin_user)): 'TOP_K': request.app.state.config.TOP_K, 'BYPASS_EMBEDDING_AND_RETRIEVAL': request.app.state.config.BYPASS_EMBEDDING_AND_RETRIEVAL, 'RAG_FULL_CONTEXT': request.app.state.config.RAG_FULL_CONTEXT, + 'RAG_MESSAGE_SCOPED_FILE_CONTEXT': request.app.state.config.RAG_MESSAGE_SCOPED_FILE_CONTEXT, # Hybrid search settings 'ENABLE_RAG_HYBRID_SEARCH': request.app.state.config.ENABLE_RAG_HYBRID_SEARCH, 'ENABLE_RAG_HYBRID_SEARCH_ENRICHED_TEXTS': request.app.state.config.ENABLE_RAG_HYBRID_SEARCH_ENRICHED_TEXTS, @@ -636,6 +637,7 @@ class ConfigForm(BaseModel): TOP_K: int | None = None BYPASS_EMBEDDING_AND_RETRIEVAL: bool | None = None RAG_FULL_CONTEXT: bool | None = None + RAG_MESSAGE_SCOPED_FILE_CONTEXT: bool | None = None # Hybrid search settings ENABLE_RAG_HYBRID_SEARCH: bool | None = None @@ -731,6 +733,11 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend if form_data.RAG_FULL_CONTEXT is not None else request.app.state.config.RAG_FULL_CONTEXT ) + request.app.state.config.RAG_MESSAGE_SCOPED_FILE_CONTEXT = ( + form_data.RAG_MESSAGE_SCOPED_FILE_CONTEXT + if form_data.RAG_MESSAGE_SCOPED_FILE_CONTEXT is not None + else request.app.state.config.RAG_MESSAGE_SCOPED_FILE_CONTEXT + ) # Hybrid search settings request.app.state.config.ENABLE_RAG_HYBRID_SEARCH = ( @@ -1131,6 +1138,7 @@ async def update_rag_config(request: Request, form_data: ConfigForm, user=Depend 'TOP_K': request.app.state.config.TOP_K, 'BYPASS_EMBEDDING_AND_RETRIEVAL': request.app.state.config.BYPASS_EMBEDDING_AND_RETRIEVAL, 'RAG_FULL_CONTEXT': request.app.state.config.RAG_FULL_CONTEXT, + 'RAG_MESSAGE_SCOPED_FILE_CONTEXT': request.app.state.config.RAG_MESSAGE_SCOPED_FILE_CONTEXT, # Hybrid search settings 'ENABLE_RAG_HYBRID_SEARCH': request.app.state.config.ENABLE_RAG_HYBRID_SEARCH, 'TOP_K_RERANKER': request.app.state.config.TOP_K_RERANKER, diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 576afe3b24..c40eebd24e 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -954,6 +954,153 @@ def get_source_context(sources: list, source_ids: dict = None, include_content: return context_string +def is_file_context_item(item: dict) -> bool: + if not isinstance(item, dict): + return False + + return item.get('type') != 'image' and not (item.get('content_type') or '').startswith('image/') + + +def get_file_context_key(item: dict) -> tuple: + if item.get('id') is not None: + return (item.get('type', 'file'), item.get('id')) + if item.get('url') is not None: + return ('url', item.get('url')) + + try: + return ('item', json.dumps(item, sort_keys=True)) + except TypeError: + return ('item', str(item)) + + +def set_message_text_content(message: dict, content: str) -> None: + if isinstance(message.get('content'), list): + for item in message['content']: + if isinstance(item, dict) and item.get('type') == 'text': + item['text'] = content + return + message['content'].insert(0, {'type': 'text', 'text': content}) + else: + message['content'] = content + + +def strip_message_files(messages: list[dict]) -> None: + for message in messages: + message.pop('files', None) + + +def attach_current_user_message_files(messages: list[dict], metadata: dict) -> None: + user_message = metadata.get('user_message') or {} + files = [item for item in user_message.get('files', []) if is_file_context_item(item)] + + if not files: + return + + for message in reversed(messages): + if message.get('role') == 'user' and not message.get('files'): + message['files'] = files + return + + +async def expand_file_context_items(items: list[dict], user: UserModel) -> list[dict]: + expanded_items = [] + + for item in items: + if item.get('type', 'file') == 'folder': + folder_id = item.get('id') + if folder_id: + folder = await Folders.get_folder_by_id_and_user_id(folder_id, user.id) + if folder and folder.data and 'files' in folder.data: + expanded_items.extend(await get_accessible_folder_files(folder.data['files'], user)) + continue + + expanded_items.append(item) + + return expanded_items + + +async def apply_message_file_contexts( + request: Request, + messages: list[dict], + files: list[dict] | None, + metadata: dict, + user: UserModel, +) -> tuple[list[dict], list[dict] | None, list[dict], list[tuple[int, str, str]]]: + attach_current_user_message_files(messages, metadata) + + handled_file_keys = set() + source_ids = {} + scoped_sources = [] + message_contexts = [] + + for message_index, message in enumerate(messages): + if message.get('role') != 'user': + continue + + message_files = [item for item in message.get('files', []) if is_file_context_item(item)] + if not message_files: + continue + + query = get_content_from_message(message) or '' + expanded_items = await expand_file_context_items(message_files, user) + + if not expanded_items: + continue + + all_full_context = all(item.get('context') == 'full' for item in expanded_items) + + try: + sources = await get_sources_from_items( + request=request, + items=expanded_items, + queries=[query], + embedding_function=lambda query, prefix: request.app.state.EMBEDDING_FUNCTION( + query, prefix=prefix, user=user + ), + k=request.app.state.config.TOP_K, + reranking_function=( + (lambda query, documents: request.app.state.RERANKING_FUNCTION(query, documents, user=user)) + if request.app.state.RERANKING_FUNCTION + else None + ), + k_reranker=request.app.state.config.TOP_K_RERANKER, + r=request.app.state.config.RELEVANCE_THRESHOLD, + hybrid_bm25_weight=request.app.state.config.HYBRID_BM25_WEIGHT, + hybrid_search=request.app.state.config.ENABLE_RAG_HYBRID_SEARCH, + full_context=all_full_context or request.app.state.config.RAG_FULL_CONTEXT, + user=user, + ) + except Exception as e: + log.exception(e) + continue + + context = get_source_context(sources, source_ids=source_ids).strip() + if context: + message_contexts.append((message_index, context, query)) + scoped_sources.extend(sources) + handled_file_keys.update(get_file_context_key(item) for item in message_files) + + if files is not None and handled_file_keys: + files = [item for item in files if get_file_context_key(item) not in handled_file_keys] + + return messages, files, scoped_sources, message_contexts + + +async def apply_message_file_context_templates( + request: Request, + messages: list[dict], + message_contexts: list[tuple[int, str, str]], +) -> list[dict]: + for message_index, context, query in message_contexts: + if 0 <= message_index < len(messages): + set_message_text_content( + messages[message_index], + await rag_template(request.app.state.config.RAG_TEMPLATE, context, query), + ) + + return messages + + async def apply_source_context_to_messages( request: Request, messages: list, @@ -2394,8 +2541,15 @@ async def process_chat_payload(request, form_data, user, metadata, model): if f.get('url') ], ] - # Strip files field — it's been incorporated into content - message.pop('files', None) + # Keep non-image files only long enough for optional message-scoped RAG. + if request.app.state.config.RAG_MESSAGE_SCOPED_FILE_CONTEXT: + non_image_files = [file for file in message.get('files', []) if is_file_context_item(file)] + if non_image_files: + message['files'] = non_image_files + else: + message.pop('files', None) + else: + message.pop('files', None) if regeneration_prompt: form_data['messages'].append({'role': 'user', 'content': regeneration_prompt}) @@ -2449,6 +2603,11 @@ async def process_chat_payload(request, form_data, user, metadata, model): events = [] sources = [] + message_scoped_sources = [] + message_scoped_contexts = [] + file_context_enabled = (model.get('info', {}).get('meta', {}).get('capabilities') or {}).get( + 'file_context', True + ) # Folder "Project" handling # Check if the request has chat_id and is inside of a folder @@ -2677,6 +2836,20 @@ async def process_chat_payload(request, form_data, user, metadata, model): # if prompt and len(prompt or "") < 500 and (not files or len(files) == 0): # urls = extract_urls(prompt) + if request.app.state.config.RAG_MESSAGE_SCOPED_FILE_CONTEXT: + if file_context_enabled: + form_data['messages'], files, message_scoped_sources, message_scoped_contexts = ( + await apply_message_file_contexts( + request, + form_data['messages'], + files, + metadata, + user, + ) + ) + + strip_message_files(form_data['messages']) + if files: if not files: files = [] @@ -2883,9 +3056,6 @@ async def process_chat_payload(request, form_data, user, metadata, model): except Exception as e: log.exception(e) - # Check if file context extraction is enabled for this model (default True) - file_context_enabled = (model.get('info', {}).get('meta', {}).get('capabilities') or {}).get('file_context', True) - if file_context_enabled: try: form_data, flags = await chat_completion_files_handler(request, form_data, extra_params, user) @@ -2893,13 +3063,21 @@ async def process_chat_payload(request, form_data, user, metadata, model): except Exception as e: log.exception(e) + if message_scoped_contexts: + form_data['messages'] = await apply_message_file_context_templates( + request, + form_data['messages'], + message_scoped_contexts, + ) + # Save the pre-RAG message state so the native tool call loop can # restore to the true original (before file-source injection) rather # than a snapshot that already has the RAG template baked in. + all_sources = [*message_scoped_sources, *sources] system_message = get_system_message(form_data['messages']) metadata['system_prompt'] = get_content_from_message(system_message) if system_message else None - metadata['user_prompt'] = get_last_user_message(form_data['messages']) - metadata['sources'] = sources[:] if sources else [] + metadata['user_prompt'] = prompt + metadata['sources'] = all_sources[:] if all_sources else [] # If context is not empty, insert it into the messages if sources and prompt: @@ -2908,7 +3086,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): # If there are citations, add them to the data_items sources = [ source - for source in sources + for source in all_sources if source.get('source', {}).get('name', '') or source.get('source', {}).get('id', '') ] diff --git a/src/lib/components/admin/Settings/Documents.svelte b/src/lib/components/admin/Settings/Documents.svelte index 92706ff94d..34e7f5afa6 100644 --- a/src/lib/components/admin/Settings/Documents.svelte +++ b/src/lib/components/admin/Settings/Documents.svelte @@ -811,6 +811,22 @@ +
+
+ + {$i18n.t('Scope File Context by Message')} + +
+
+ +
+
+ {#if !RAGConfig.BYPASS_EMBEDDING_AND_RETRIEVAL}
{$i18n.t('Text Splitter')}