diff --git a/backend/open_webui/retrieval/external.py b/backend/open_webui/retrieval/external.py index dcfc575fc8..1f0c06d8a7 100644 --- a/backend/open_webui/retrieval/external.py +++ b/backend/open_webui/retrieval/external.py @@ -4,6 +4,7 @@ import re import time from typing import Any, Optional +from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX from open_webui.models.config import Config from open_webui.models.knowledge import KnowledgeModel @@ -103,7 +104,7 @@ async def _retrieve_qdrant(connection, auth_config, knowledge, query, count, emb source_config = _source_config(knowledge) vector_field = source_config.get('vector_field') or None - vector = await embedding_function(query) + vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX) def _search(): client = QdrantClient( @@ -152,7 +153,7 @@ async def _retrieve_milvus(connection, auth_config, knowledge, query, count, emb content_field = source_config.get('content_field') or 'data.text' metadata_field = source_config.get('metadata_field') or 'metadata' - vector = await embedding_function(query) + vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX) def _search(): client_kwargs = { @@ -229,7 +230,7 @@ async def _retrieve_pgvector(connection, auth_config, knowledge, query, count, e metadata_field = source_config.get('metadata_field') or 'vmetadata' document_id_field = source_config.get('document_id_field') or 'id' - vector = await embedding_function(query) + vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX) def _search(): from psycopg import sql diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index db3ba0b86d..b80e448149 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -12,7 +12,7 @@ from urllib.parse import quote from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from fastapi.responses import StreamingResponse -from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL +from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, RAG_EMBEDDING_CONTENT_PREFIX from open_webui.constants import ERROR_MESSAGES from open_webui.events import EVENTS, publish_event from open_webui.internal.db import get_async_session @@ -74,7 +74,7 @@ async def embed_knowledge_base_metadata( """Generate and store embedding for knowledge base.""" try: content = f'{name}\n\n{description}' if description else name - embedding = await request.app.state.EMBEDDING_FUNCTION(content) + embedding = await request.app.state.EMBEDDING_FUNCTION(content, prefix=RAG_EMBEDDING_CONTENT_PREFIX) await ASYNC_VECTOR_DB_CLIENT.upsert( collection_name=KNOWLEDGE_BASES_COLLECTION, items=[ diff --git a/backend/open_webui/routers/memories.py b/backend/open_webui/routers/memories.py index 8e6d50f39a..26c8293d88 100644 --- a/backend/open_webui/routers/memories.py +++ b/backend/open_webui/routers/memories.py @@ -5,13 +5,13 @@ import logging from typing import Literal from fastapi import APIRouter, Depends, HTTPException, Request, status +from open_webui.config import RAG_EMBEDDING_CONTENT_PREFIX, RAG_EMBEDDING_QUERY_PREFIX from open_webui.constants import ERROR_MESSAGES from open_webui.events import EVENTS, publish_event from open_webui.internal.db import get_async_session from open_webui.models.config import Config from open_webui.models.memories import Memories, MemoryModel from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT -from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX from open_webui.utils.access_control import has_permission from open_webui.utils.auth import get_verified_user from open_webui.utils.memory import ( @@ -148,7 +148,9 @@ async def add_memory( meta={'created_by': 'manual'}, ) - vector = await request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user) + vector = await request.app.state.EMBEDDING_FUNCTION( + memory_vector_text(memory.content, memory.path), prefix=RAG_EMBEDDING_CONTENT_PREFIX, user=user + ) await ASYNC_VECTOR_DB_CLIENT.upsert( collection_name=f'user-memory-{user.id}', @@ -208,6 +210,7 @@ async def update_memories( if result.get('status') in {'created', 'updated'}: vector = await request.app.state.EMBEDDING_FUNCTION( memory_vector_text(memory.content, memory.path), + prefix=RAG_EMBEDDING_CONTENT_PREFIX, user=user, ) upsert_items.append( @@ -284,7 +287,7 @@ async def query_memory( if not memories: raise HTTPException(status_code=404, detail='No memories found for user') - vector = await request.app.state.EMBEDDING_FUNCTION(form_data.content, RAG_EMBEDDING_QUERY_PREFIX, user=user) + vector = await request.app.state.EMBEDDING_FUNCTION(form_data.content, prefix=RAG_EMBEDDING_QUERY_PREFIX, user=user) results = await ASYNC_VECTOR_DB_CLIENT.search( collection_name=f'user-memory-{user.id}', @@ -407,7 +410,9 @@ async def reset_memory_from_vector_db( # Generate vectors in parallel vectors = await asyncio.gather( *[ - request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user) + request.app.state.EMBEDDING_FUNCTION( + memory_vector_text(memory.content, memory.path), prefix=RAG_EMBEDDING_CONTENT_PREFIX, user=user + ) for memory in memories ] ) @@ -503,7 +508,9 @@ async def update_memory_by_id( raise HTTPException(status_code=404, detail=ERROR_MESSAGES.NOT_FOUND) if form_data.content is not None or form_data.path is not None: - vector = await request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user) + vector = await request.app.state.EMBEDDING_FUNCTION( + memory_vector_text(memory.content, memory.path), prefix=RAG_EMBEDDING_CONTENT_PREFIX, user=user + ) await ASYNC_VECTOR_DB_CLIENT.upsert( collection_name=f'user-memory-{user.id}', diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index 208f3e7dbe..5d8754291a 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -16,6 +16,7 @@ from typing import Literal, Optional from fastapi import HTTPException, Request +from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX from open_webui.models.channels import Channel, ChannelMember, Channels from open_webui.models.chats import Chats from open_webui.models.config import Config @@ -3237,7 +3238,7 @@ async def query_knowledge_bases( user_id = __user__.get('id') user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] - query_embedding = await __request__.app.state.EMBEDDING_FUNCTION(query) + query_embedding = await __request__.app.state.EMBEDDING_FUNCTION(query, prefix=RAG_EMBEDDING_QUERY_PREFIX) # Min-heap of (distance, knowledge_base_id) - only holds top `count` results top_results_heap = []