mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-13 01:02:25 -06:00
refac
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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=[
|
||||
|
||||
@@ -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}',
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user