diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index 58a4da1f1b..1d2fdfd01a 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -703,6 +703,7 @@ async def query_collection( queries: list[str], embedding_function, k: int, + user: UserModel | None = None, ) -> dict: config = await Config.get_many( 'rag.enable_hybrid_search', @@ -715,7 +716,7 @@ async def query_collection( if request and config.get('rag.enable_hybrid_search'): try: reranking_function = ( - (lambda query, documents: request.app.state.RERANKING_FUNCTION(query, documents)) + (lambda query, documents: request.app.state.RERANKING_FUNCTION(query, documents, user=user)) if request.app.state.RERANKING_FUNCTION else None ) @@ -1692,6 +1693,7 @@ async def get_sources_from_items( queries=queries, embedding_function=embedding_function, k=k, + user=user, ) except Exception as e: log.exception(e) diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index e048098bae..226672ec9d 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -3221,6 +3221,7 @@ async def query_collection_handler( query, prefix=prefix, user=user ), k=form_data.k if form_data.k else config.TOP_K, + user=user, ) except HTTPException: diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index ed4974dec9..5943df6d39 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -3316,6 +3316,7 @@ async def query_knowledge_files( queries=[query], embedding_function=lambda queries, prefix: embedding_function(queries, prefix=prefix, user=user_model), k=count, + user=user_model, ) if query_results and 'documents' in query_results: