mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-06 02:48:04 +00:00
Add message-scoped file context option
This commit is contained in:
parent
f85cb27ef8
commit
f7a6d517af
5 changed files with 218 additions and 8 deletions
|
|
@ -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',
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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', '')
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -811,6 +811,22 @@
|
|||
</div>
|
||||
</div>
|
||||
|
||||
<div class=" mb-2.5 flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">
|
||||
<Tooltip
|
||||
content={$i18n.t(
|
||||
'Keep attached files scoped to the user message they were sent with instead of merging all chat files into the latest RAG context.'
|
||||
)}
|
||||
placement="top-start"
|
||||
>
|
||||
{$i18n.t('Scope File Context by Message')}
|
||||
</Tooltip>
|
||||
</div>
|
||||
<div class="flex items-center relative">
|
||||
<Switch bind:state={RAGConfig.RAG_MESSAGE_SCOPED_FILE_CONTEXT} />
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{#if !RAGConfig.BYPASS_EMBEDDING_AND_RETRIEVAL}
|
||||
<div class=" mb-2.5 flex w-full justify-between">
|
||||
<div class=" self-center text-xs font-medium">{$i18n.t('Text Splitter')}</div>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue